@@ -245,7 +245,13 @@ def state_dict(self):
245245
246246
247247class MixedDataset (IterableDataset , Stateful ):
248- def __init__ (self , datasets : list [IterableDataset ], weights : list [float ] | None ):
248+ def __init__ (
249+ self ,
250+ datasets : list [IterableDataset ],
251+ dp_rank : int ,
252+ weights : list [float ] | None ,
253+ seed : int | None = 0 ,
254+ ):
249255 self .datasets = datasets
250256
251257 _initial_weights = [1.0 ] * len (self .datasets ) if weights is None else weights
@@ -260,7 +266,7 @@ def __init__(self, datasets: list[IterableDataset], weights: list[float] | None)
260266 self ._dataset_indices = list (range (len (self .datasets )))
261267 self ._sample_idx = 0
262268 self ._data_iters = None
263- self ._rng = Random (self . _sample_idx )
269+ self ._rng = Random (seed + dp_rank )
264270
265271 @property
266272 def normed_weights (self ):
@@ -271,7 +277,6 @@ def _init_data_iters(self):
271277 self ._data_iters = [iter (dataset ) for dataset in self .datasets ]
272278
273279 def _sample_dataset (self , sample_idx : int ):
274- self ._rng .seed (sample_idx )
275280 dataset_index = self ._rng .choices (
276281 self ._dataset_indices , weights = self .weights .tolist ()
277282 )[0 ]
@@ -330,11 +335,11 @@ def load_state_dict(self, state_dict):
330335 self .weights .copy_ (torch .tensor (loaded_weights , dtype = torch .float64 ))
331336 self .num_sampled_per_dataset .copy_ (state_dict ["num_sampled_per_dataset" ])
332337
338+ self ._rng .setstate (state_dict ["rng_state" ])
333339 # Restore sub-datasets.
334340 dataset_dicts = state_dict ["datasets" ]
335341 for dataset in self .datasets :
336342 dataset .load_state_dict (dataset_dicts [dataset .dataset_name ])
337-
338343 # Unset data iterators so they will be re-initialized.
339344 self ._data_iters = None
340345
@@ -346,6 +351,7 @@ def state_dict(self):
346351 "datasets" : {
347352 dataset .dataset_name : dataset .state_dict () for dataset in self .datasets
348353 },
354+ "rng_state" : self ._rng .getstate (),
349355 }
350356
351357
@@ -439,6 +445,7 @@ def __init__(
439445 self ,
440446 dataset : IterableDataset ,
441447 * ,
448+ dp_rank : int ,
442449 buffer_size : int = 10000 ,
443450 seed : int | None = 0 ,
444451 ) -> None :
@@ -448,7 +455,7 @@ def __init__(
448455 self .buffer_size = buffer_size
449456 self ._enabled = True
450457 self ._initial_seed = seed
451- self ._rng = Random (self ._initial_seed )
458+ self ._rng = Random (self ._initial_seed + dp_rank )
452459
453460 def set_shuffle (self , shuffle : bool = True ):
454461 self ._enabled = shuffle
@@ -539,8 +546,7 @@ def build_text_dataloader(
539546 rng = torch .Generator ()
540547 dataset_streaming = job_config .training .dataset_streaming
541548
542- if job_config .training .dataset_seed is not None :
543- rng .manual_seed (job_config .training .dataset_seed )
549+ rng .manual_seed (job_config .debug .seed )
544550
545551 num_mtp_tokens = job_config .training .num_mtp_tokens
546552 dataset_weights = job_config .training .dataset_weights
@@ -610,7 +616,9 @@ def build_text_dataloader(
610616
611617 # First pack, then mix → data is only mixed in batch dimension.
612618 # First mix, then pack → data is also mixed inside packed sample.
613- hf_ds = MixedDataset (hf_datasets , dataset_weights )
619+ hf_ds = MixedDataset (
620+ hf_datasets , dp_rank , dataset_weights , seed = job_config .debug .seed
621+ )
614622 if dataset_mix_in_seq :
615623 hf_ds = GreedyPackedDataset (
616624 dataset = hf_ds ,
@@ -619,14 +627,12 @@ def build_text_dataloader(
619627 num_mtp_tokens = num_mtp_tokens ,
620628 )
621629
622- if job_config .training .dataset_seed is None :
623- job_config .training .dataset_seed = job_config .debug .seed
624-
625630 if job_config .training .dataset_shuffle_buffer_size :
626631 hf_ds = WindowShuffledDataset (
627632 hf_ds ,
633+ dp_rank = dp_rank ,
628634 buffer_size = job_config .training .dataset_shuffle_buffer_size ,
629- seed = job_config .training . dataset_seed ,
635+ seed = job_config .debug . seed ,
630636 )
631637 prefetch_factor = job_config .training .dataset_prefetch_factor
632638 num_workers = job_config .training .dataset_num_workers
0 commit comments