Skip to content

Commit 93ac435

Browse files
committed
improve and fix seed for data mix
1 parent dccec36 commit 93ac435

2 files changed

Lines changed: 18 additions & 18 deletions

File tree

torchtitan/config/job_config.py

Lines changed: 0 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -372,12 +372,6 @@ class Training:
372372
dataset_pin_memory: bool = False
373373
"""Whether to use memory pinning in the data loader"""
374374

375-
dataset_seed: int | None = None
376-
"""
377-
Choose the base RNG seed used for data shuffling. By default,
378-
use the same as `training.seed`.
379-
"""
380-
381375
dataset_shuffle_buffer_size: int = 0
382376
"""Buffer size of windowed shuffling buffer. 0 means no shuffling (the default)."""
383377

torchtitan/hf_datasets/text_datasets.py

Lines changed: 18 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -245,7 +245,13 @@ def state_dict(self):
245245

246246

247247
class 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

Comments
 (0)