@@ -577,8 +577,12 @@ def copy_path(src_path: str, tgt_path: str):
577577 else :
578578 raise ValueError (f'Source path is neither a file nor a directory: { src_path } ' )
579579
580- def _prepare_data_iterator (self , train_dataset , val_dataset = None , use_origin_cyclic : bool = False ):
581- train_dataloader , val_dataloader = self ._prepare_dataloader (train_dataset , val_dataset )
580+ def _prepare_data_iterator (self ,
581+ train_dataset ,
582+ val_dataset = None ,
583+ use_origin_cyclic : bool = False ,
584+ seed : Optional [int ] = None ):
585+ train_dataloader , val_dataloader = self ._prepare_dataloader (train_dataset , val_dataset , seed = seed )
582586 train_data_iterator = iter (self .cyclic_iter (train_dataloader , use_origin_cyclic = use_origin_cyclic ))
583587 val_data_iterator = None
584588 if val_dataset is not None :
@@ -973,11 +977,15 @@ def _aggregated_metrics(self, metrics, total_metrics):
973977 total_metrics [key ] = torch .tensor ([0.0 , 0.0 ], dtype = torch .float32 , device = torch .cuda .current_device ())
974978 total_metrics [key ] += val
975979
976- def _prepare_dataloader (self , train_dataset , val_dataset = None ):
980+ def _prepare_dataloader (self , train_dataset , val_dataset = None , seed : Optional [ int ] = None ):
977981 args = self .args
978982 val_dataloader = None
983+ generator = None
984+ if seed is not None :
985+ generator = torch .Generator ()
986+ generator .manual_seed (seed )
979987 if args .streaming :
980- train_dataloader = build_streaming_dataloader (args , train_dataset , self .data_collator )
988+ train_dataloader = build_streaming_dataloader (args , train_dataset , self .data_collator , generator = generator )
981989 if val_dataset is not None :
982990 val_dataloader = build_streaming_dataloader (args , val_dataset , self .data_collator )
983991 return train_dataloader , val_dataloader
@@ -991,8 +999,9 @@ def _prepare_dataloader(self, train_dataset, val_dataset=None):
991999 data_sharding = args .data_sharding ,
9921000 shuffle = args .train_dataloader_shuffle ,
9931001 group_by_length = args .group_by_length ,
1002+ seed = seed or 0 ,
9941003 )
995- train_dataloader = self ._create_dataloader (train_dataset , train_batch_sampler )
1004+ train_dataloader = self ._create_dataloader (train_dataset , train_batch_sampler , generator = generator )
9961005 if val_dataset is not None :
9971006 val_batch_sampler = MegatronPretrainingSampler (
9981007 total_samples = len (val_dataset ),
@@ -1004,7 +1013,7 @@ def _prepare_dataloader(self, train_dataset, val_dataset=None):
10041013 val_dataloader = self ._create_dataloader (val_dataset , val_batch_sampler )
10051014 return train_dataloader , val_dataloader
10061015
1007- def _create_dataloader (self , dataset , batch_sampler ):
1016+ def _create_dataloader (self , dataset , batch_sampler , generator = None ):
10081017 args = self .args
10091018
10101019 dataloader = torch .utils .data .DataLoader (
@@ -1015,6 +1024,7 @@ def _create_dataloader(self, dataset, batch_sampler):
10151024 persistent_workers = args .dataloader_persistent_workers if args .dataloader_num_workers > 0 else False ,
10161025 prefetch_factor = args .dataloader_prefetch_factor if args .dataloader_num_workers > 0 else None ,
10171026 collate_fn = self .data_collator ,
1027+ generator = generator ,
10181028 )
10191029 return dataloader
10201030
0 commit comments