From d85779bee3e6fe11ef2f85c095c6f7b6e03034fb Mon Sep 17 00:00:00 2001 From: Isaac Hernandez Date: Wed, 5 Aug 2026 00:56:07 -0400 Subject: [PATCH] Honor seed=0 in iterate_batches iterate_batches gated its seeding on `if seed:`, so an explicit seed of 0 was silently ignored and batch order fell back to whatever state global numpy happened to be in. 0 is the default seed in mlx_lm.lora's CONFIG_DEFAULTS, which made the most common value the one that did not work. Use `if seed is not None:` instead. Add a regression test asserting a given seed produces the same batch order regardless of other numpy use in between, for both 0 and a non-zero seed. The test fails before this change. --- mlx_lm/tuner/trainer.py | 2 +- tests/test_tuner_trainer.py | 21 +++++++++++++++++++++ 2 files changed, 22 insertions(+), 1 deletion(-) 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()