diff --git a/mlx_lm/tuner/trainer.py b/mlx_lm/tuner/trainer.py index 77b7bcdbc..7ed8bbb11 100644 --- a/mlx_lm/tuner/trainer.py +++ b/mlx_lm/tuner/trainer.py @@ -135,7 +135,7 @@ def iterate_batches( idx[i + offset : i + offset + batch_size : step] for i in range(0, len(idx) - batch_size + 1, batch_size) ] - if seed: + if seed is not None: np.random.seed(seed) while True: indices = np.random.permutation(len(batch_idx)) diff --git a/tests/test_tuner_trainer.py b/tests/test_tuner_trainer.py index 09fe57029..a0bd96e26 100644 --- a/tests/test_tuner_trainer.py +++ b/tests/test_tuner_trainer.py @@ -3,6 +3,7 @@ import unittest import mlx.core as mx +import numpy as np from mlx_lm.tuner.trainer import iterate_batches @@ -49,6 +50,26 @@ def run(rank, size, batch): run(2, 4, 8) run(3, 4, 8) + def test_iterate_batches_seed(self): + # One distinct token id per row so the batch order is observable. + data = [([i + 1] * 8, 0) for i in range(64)] + + def order(seed, consume=0): + np.random.seed(1) + if consume: + # Stand-in for anything else drawing from numpy in between. + np.random.rand(consume) + batches = iterate_batches(data, 4, 8, loop=True, seed=seed) + return [b[0].tolist()[0] for b, _ in zip(batches, range(5))] + + # seed=0 must be honored like any other seed. It is also the default + # in mlx_lm.lora's CONFIG_DEFAULTS, so `if seed:` silently dropped it. + for seed in (0, 42): + with self.subTest(seed=seed): + self.assertEqual(order(seed), order(seed, consume=3)) + + self.assertNotEqual(order(0), order(42)) + if __name__ == "__main__": unittest.main()