From 56051b4d59f854e00d701e3514fabce0d728f254 Mon Sep 17 00:00:00 2001 From: alanhuangyoo Date: Tue, 8 Sep 2026 20:06:24 +0800 Subject: [PATCH] Checkpoint the RNG the curriculum sampler actually draws from Fixes #8405. Every draw in DeepSpeedDataSampler goes through self.np_rng, a np.random.default_rng seeded in __init__: self.np_rng.shuffle(new_cluster) # get_new_cluster self.np_rng.choice(num_clusters, ...) # sample_from_clusters self.np_rng.shuffle(cluster) # reshuffle_clusters self.np_rng.shuffle(batch) state_dict() stored np.random.get_state() and load_state_dict() restored it. That is the process-wide legacy RandomState, which this sampler never touches, so the saved value reads the same whether the sampler has drawn nothing or a thousand batches. On resume __init__ re-seeds self.np_rng, load_state_dict leaves it alone, and the sampler restarts its own stream: the same cluster mix per step and the same shuffles it produced the first time. Restoring the global state also moved whatever else in the process draws from np.random. Save and restore self.np_rng.bit_generator.state instead. A checkpoint written before this holds the legacy tuple, which carries no sampler stream to restore; _load_np_rng_state recognises it, warns once on rank 0, and leaves the sampler on its configured seed -- where those checkpoints already resumed -- rather than refusing to load. tests/unit/runtime/data_pipeline/test_curriculum_sampler_rng.py: the saved state moves as the sampler draws, a resumed sampler continues the stream instead of replaying it (through load_state_dict, not by poking the generator), loading leaves the global numpy RNG alone, and a legacy checkpoint still loads. Two of the four fail on master. Reported by @ebarkhordar. Signed-off-by: alanhuangyoo --- .../data_sampling/data_sampler.py | 26 +++- .../test_curriculum_sampler_rng.py | 123 ++++++++++++++++++ 2 files changed, 147 insertions(+), 2 deletions(-) create mode 100644 tests/unit/runtime/data_pipeline/test_curriculum_sampler_rng.py diff --git a/deepspeed/runtime/data_pipeline/data_sampling/data_sampler.py b/deepspeed/runtime/data_pipeline/data_sampling/data_sampler.py index 100bef3f7946..24cf30e6ec49 100644 --- a/deepspeed/runtime/data_pipeline/data_sampling/data_sampler.py +++ b/deepspeed/runtime/data_pipeline/data_sampling/data_sampler.py @@ -313,6 +313,22 @@ def __iter__(self): self.consumed_samples += len(current_batch) current_batch = [] + def _load_np_rng_state(self, rng_state): + """Restore `self.np_rng`, tolerating checkpoints written before this was fixed. + + Older checkpoints hold the tuple `np.random.get_state()` returns for the + legacy global RandomState. There is no sampler stream recorded in it, so the + best that can be done is to leave `self.np_rng` on its fresh seed -- the same + position those checkpoints already resumed at -- rather than fail to load. + """ + if isinstance(rng_state, dict): + self.np_rng.bit_generator.state = rng_state + return + if self.global_rank == 0: + logger.warning("Curriculum learning checkpoint stores the legacy global numpy RNG state, " + "which does not describe this sampler's own stream. Resuming its sampling " + "from the configured seed.") + def state_dict(self): return { CURRICULUM_LEARNING_BATCH: self.batch, @@ -321,7 +337,13 @@ def state_dict(self): CURRICULUM_LEARNING_CURRENT_DIFFICULTIES: self.current_difficulties, CURRICULUM_LEARNING_DATA_CLUSTER_PATHS: self.data_cluster_paths, CURRICULUM_LEARNING_DATA_CLUSTER_CURRENT_POSITION: self.data_cluster_current_position, - CURRICULUM_LEARNING_NP_RNG_STATE: np.random.get_state() + # `self.np_rng` is what every draw in this class goes through + # (sample_from_clusters, get_new_cluster, reshuffle_clusters), so its + # bit generator's state is the one worth saving. `np.random.get_state()` + # returns the process-wide legacy RandomState, which this sampler never + # touches: it reads the same whether the sampler has drawn nothing or a + # thousand batches. + CURRICULUM_LEARNING_NP_RNG_STATE: self.np_rng.bit_generator.state } def load_state_dict(self, state_dict): @@ -331,7 +353,7 @@ def load_state_dict(self, state_dict): self.current_difficulties = state_dict[CURRICULUM_LEARNING_CURRENT_DIFFICULTIES] self.data_cluster_paths = state_dict[CURRICULUM_LEARNING_DATA_CLUSTER_PATHS] self.data_cluster_current_position = state_dict[CURRICULUM_LEARNING_DATA_CLUSTER_CURRENT_POSITION] - np.random.set_state(state_dict[CURRICULUM_LEARNING_NP_RNG_STATE]) + self._load_np_rng_state(state_dict[CURRICULUM_LEARNING_NP_RNG_STATE]) cluster_root_path = self.data_efficiency_config[DATA_SAMPLING][CURRICULUM_LEARNING][ CURRICULUM_LEARNING_CLUSTER_PATH] # Backward compatibility: previously data_cluster_paths were stored as diff --git a/tests/unit/runtime/data_pipeline/test_curriculum_sampler_rng.py b/tests/unit/runtime/data_pipeline/test_curriculum_sampler_rng.py new file mode 100644 index 000000000000..02cfced6e0b6 --- /dev/null +++ b/tests/unit/runtime/data_pipeline/test_curriculum_sampler_rng.py @@ -0,0 +1,123 @@ +# Copyright (c) DeepSpeed Team. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team +"""The curriculum sampler must checkpoint the generator it actually draws from. + +Every draw in DeepSpeedDataSampler goes through `self.np_rng`, a +`np.random.default_rng`. `state_dict` used to store `np.random.get_state()` -- +the process-wide legacy RandomState, which the sampler never touches -- so the +saved value carried no information about how far the sampling stream had run, +and resume replayed it from the seed. +""" + +import hashlib + +import numpy as np + +from deepspeed.runtime.data_pipeline.config import get_data_efficiency_config +from deepspeed.runtime.data_pipeline.constants import CURRICULUM_LEARNING_NP_RNG_STATE +from deepspeed.runtime.data_pipeline.data_sampling.data_sampler import DeepSpeedDataSampler + +CLUSTER_SIZES = [10, 20, 30, 40] + + +def _config(): + metric = { + "index_to_sample_path": "dummy", + "index_to_metric_path": "dummy", + "difficulty_type": "value", + "clustering_type": "single_cluster", + "min_difficulty": 8, + "max_difficulty": 80, + "schedule_type": "fixed_linear", + "schedule_config": { + "total_curriculum_step": 100, + "difficulty_step": 8 + }, + } + return get_data_efficiency_config({ + "data_efficiency": { + "enabled": True, + "seed": 1234, + "data_sampling": { + "enabled": True, + "curriculum_learning": { + "enabled": True, + "data_cluster_path": "/tmp/clusters", + "curriculum_metrics": { + "dummy": metric + }, + }, + }, + } + }) + + +def _sampler(): + sampler = DeepSpeedDataSampler(_config(), 100, 8, 0, 1, None, 1, global_rank=0) + sampler.data_clusters = [None] * len(CLUSTER_SIZES) + sampler.data_cluster_sizes = list(CLUSTER_SIZES) + return sampler + + +def _state_fingerprint(sampler): + """A short, comparable digest of the checkpointed RNG state. + + Hashed rather than compared directly so this reads the same whether the + checkpoint holds a bit-generator dict or the legacy RandomState tuple, and so a + failure prints a digest instead of 624 words of Mersenne Twister. + """ + return hashlib.sha256(repr(sampler.state_dict()[CURRICULUM_LEARNING_NP_RNG_STATE]).encode()).hexdigest()[:16] + + +def test_saved_state_moves_as_the_sampler_draws(): + sampler = _sampler() + before = _state_fingerprint(sampler) + for _ in range(3): + sampler.sample_from_clusters() + after = _state_fingerprint(sampler) + + assert before != after, "the saved RNG state must reflect the draws the sampler made" + + +def test_resume_continues_the_stream_rather_than_replaying_it(): + saved = _sampler() + for _ in range(3): + saved.sample_from_clusters() + # data_cluster_paths empty so load_state_dict does no file I/O; everything else + # goes through the real API. + checkpoint = dict(saved.state_dict(), data_cluster_paths=[]) + expected = [saved.sample_from_clusters().tolist() for _ in range(3)] + + resumed = _sampler() + resumed.load_state_dict(checkpoint) + actual = [resumed.sample_from_clusters().tolist() for _ in range(3)] + + assert actual == expected, "resume must continue the sampling stream, not restart it" + + replayed = _sampler() + assert [replayed.sample_from_clusters().tolist() for _ in range(3)] != expected, \ + "a sampler on a fresh seed must not already agree -- otherwise this proves nothing" + + +def test_the_global_numpy_rng_is_left_alone(): + sampler = _sampler() + for _ in range(3): + sampler.sample_from_clusters() + + np.random.seed(0) + before = np.random.get_state()[1].copy() + sampler.load_state_dict(dict(sampler.state_dict(), data_cluster_paths=[])) + after = np.random.get_state()[1] + + assert np.array_equal(before, after), "loading must not move the process-wide numpy RNG" + + +def test_a_legacy_checkpoint_still_loads(): + # Written before this was fixed: a tuple from np.random.get_state(). + sampler = _sampler() + legacy = dict(sampler.state_dict(), data_cluster_paths=[]) + legacy[CURRICULUM_LEARNING_NP_RNG_STATE] = np.random.get_state() + + sampler.load_state_dict(legacy) # must not raise