Skip to content

Commit 8499680

Browse files
authored
Use test data zipfile in CI (#662)
* make fewer requests in tests * token
1 parent 99f8fd6 commit 8499680

4 files changed

Lines changed: 14 additions & 9 deletions

File tree

.github/workflows/pull_request.yml

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -38,4 +38,6 @@ jobs:
3838
- name: Run tests
3939
shell: bash -l {0}
4040
run: |
41-
python -m xmlrunner discover -v tests -o test-results/
41+
curl -o test_data.zip -L https://github.com/ml-explore/mlx-lm/releases/download/test_data/test_data.zip
42+
unzip test_data.zip
43+
HF_HOME="." python -m xmlrunner discover -v tests -o test-results/

tests/test_datsets.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,8 @@ def setUpClass(cls):
2121
cls.test_dir = cls.test_dir_fid.name
2222
if not os.path.isdir(cls.test_dir):
2323
os.mkdir(cls.test_dir_fid.name)
24+
# Only one HF request
25+
AutoTokenizer.from_pretrained(HF_MODEL_PATH)
2426

2527
@classmethod
2628
def tearDownClass(cls):
@@ -37,7 +39,7 @@ def test_text(self):
3739
data = {"text": "This is an example for the model."}
3840
self.save_data(4 * [data])
3941
args = types.SimpleNamespace(train=True, test=False, data=self.test_dir)
40-
tokenizer = AutoTokenizer.from_pretrained(HF_MODEL_PATH)
42+
tokenizer = AutoTokenizer.from_pretrained(HF_MODEL_PATH, local_files_only=True)
4143
train, valid, test = datasets.load_dataset(args, tokenizer)
4244
self.assertEqual(len(train), 4)
4345
self.assertEqual(len(valid), 4)
@@ -50,7 +52,7 @@ def test_completions(self):
5052
data = {"prompt": "What is the capital of France?", "completion": "Paris."}
5153
self.save_data(4 * [data])
5254
args = types.SimpleNamespace(train=True, test=False, data=self.test_dir)
53-
tokenizer = AutoTokenizer.from_pretrained(HF_MODEL_PATH)
55+
tokenizer = AutoTokenizer.from_pretrained(HF_MODEL_PATH, local_files_only=True)
5456
train, valid, test = datasets.load_dataset(args, tokenizer)
5557
self.assertEqual(len(train), 4)
5658
self.assertEqual(len(valid), 4)
@@ -69,7 +71,7 @@ def test_chat(self):
6971
}
7072
self.save_data(4 * [data])
7173
args = types.SimpleNamespace(train=True, test=False, data=self.test_dir)
72-
tokenizer = AutoTokenizer.from_pretrained(HF_MODEL_PATH)
74+
tokenizer = AutoTokenizer.from_pretrained(HF_MODEL_PATH, local_files_only=True)
7375
train, valid, test = datasets.load_dataset(args, tokenizer)
7476
self.assertEqual(len(train), 4)
7577
self.assertEqual(len(valid), 4)
@@ -91,7 +93,7 @@ def test_hf(self):
9193
test=False,
9294
train=True,
9395
)
94-
tokenizer = AutoTokenizer.from_pretrained(HF_MODEL_PATH)
96+
tokenizer = AutoTokenizer.from_pretrained(HF_MODEL_PATH, local_files_only=True)
9597
train, valid, test = datasets.load_dataset(args, tokenizer)
9698
self.assertTrue(len(train) > 0)
9799
self.assertTrue(len(train[0]) > 0)

tests/test_generate.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -81,7 +81,7 @@ def logits_processor(toks, logits):
8181

8282
def test_stream_generate_speculative(self):
8383
# Use same model as draft model, this is not a speed test
84-
draft_model, _ = load(self.HF_MODEL_PATH)
84+
draft_model = self.model
8585

8686
results: List[GenerationResponse] = []
8787
drafted: List[bool] = []

tests/test_prompt_cache.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@ class TestPromptCache(unittest.TestCase):
3434
def setUpClass(cls):
3535
cls.test_dir_fid = tempfile.TemporaryDirectory()
3636
cls.test_dir = cls.test_dir_fid.name
37+
cls.model, cls.tokenizer = load(HF_MODEL_PATH)
3738

3839
@classmethod
3940
def tearDownClass(cls):
@@ -132,7 +133,7 @@ def test_save_load_mixed_cache(self):
132133
self.assertTrue(mx.array_equal(v, lv))
133134

134135
def test_cache_with_generate(self):
135-
model, tokenizer = load(HF_MODEL_PATH)
136+
model, tokenizer = self.model, self.tokenizer
136137
prompt = tokenizer.encode("this is a prompt", return_tensors="mlx")[0]
137138
results = list(generate_step(prompt, model, max_tokens=4))
138139
toks, all_logits = zip(*results)
@@ -212,7 +213,7 @@ def test_trim_cache(self):
212213
self.assertEqual(num_trimmed, 3)
213214

214215
def test_trim_cache_with_generate(self):
215-
model, tokenizer = load(HF_MODEL_PATH)
216+
model, tokenizer = self.model, self.tokenizer
216217
prompt = tokenizer.encode("this is a prompt", return_tensors="mlx")[0]
217218

218219
prompt_cache = make_prompt_cache(model)
@@ -289,7 +290,7 @@ def test_save_load_quantized_cache(self):
289290
self.assertEqual(metadata, loaded_metadata)
290291

291292
def test_cache_to_quantized(self):
292-
model, tokenizer = load(HF_MODEL_PATH)
293+
model, tokenizer = self.model, self.tokenizer
293294
prompt = tokenizer.encode("this is a prompt", return_tensors="mlx")[0]
294295
results = zip(range(4), generate_step(prompt, model))
295296
toks, all_logits = zip(*(r[1] for r in results))

0 commit comments

Comments
 (0)