@@ -13,12 +13,13 @@ def get_tokenizer(model_name: str = "gpt2"):
1313 return tokenizer
1414
1515class StreamingDataset (IterableDataset ):
16- def __init__ (self , dataset , tokenizer , max_seq_len , mode = "pretrain" , buffer_size = 500 ):
16+ def __init__ (self , dataset , tokenizer , max_seq_len , mode = "pretrain" , buffer_size = 500 , skip_samples = 0 ):
1717 self .dataset = dataset
1818 self .tokenizer = tokenizer
1919 self .max_seq_len = max_seq_len
2020 self .mode = mode
2121 self .buffer_size = buffer_size
22+ self .skip_samples = skip_samples
2223
2324 def _prepare_sft_text (self , example ):
2425 if 'messages' in example :
@@ -39,7 +40,10 @@ def _prepare_sft_text(self, example):
3940 def __iter__ (self ) -> Iterator [Dict [str , torch .Tensor ]]:
4041 iterator = iter (self .dataset )
4142 buffer = []
42-
43+
44+ # Calculate roughly how many items to skip if they were yielded
45+ # We process skipping in the yield loop
46+
4347 for example in iterator :
4448 text = (example .get ('text' , '' ) if self .mode == "pretrain"
4549 else self ._prepare_sft_text (example ))
@@ -70,16 +74,32 @@ def __iter__(self) -> Iterator[Dict[str, torch.Tensor]]:
7074 if len (buffer ) >= self .buffer_size :
7175 random .shuffle (buffer )
7276 for _ in range (self .buffer_size // 2 ):
73- yield buffer .pop ()
77+ item = buffer .pop ()
78+ if self .skip_samples > 0 :
79+ self .skip_samples -= 1
80+ continue
81+ yield item
7482
7583 # Yield remaining
7684 random .shuffle (buffer )
7785 while buffer :
78- yield buffer .pop ()
86+ item = buffer .pop ()
87+ if self .skip_samples > 0 :
88+ self .skip_samples -= 1
89+ continue
90+ yield item
7991
80- def create_streaming_loader (dataset_name , split , tokenizer , config , batch_size , mode = "pretrain" , hf_token = None ):
92+ def create_streaming_loader (dataset_name , split , tokenizer , config , batch_size , mode = "pretrain" , hf_token = None , start_step = 0 ):
8193 raw_dataset = load_dataset (dataset_name , split = split , streaming = True ,
8294 trust_remote_code = True , token = hf_token )
83- stream_ds = StreamingDataset (raw_dataset , tokenizer , config .max_seq_len , mode = mode )
95+
96+ # Calculate samples to skip: start_step * batch_size
97+ skip_samples = start_step * batch_size
98+ if skip_samples > 0 :
99+ print (f" [Loader] Resuming: Fast-forwarding dataset by { skip_samples } samples..." )
100+
101+ stream_ds = StreamingDataset (raw_dataset , tokenizer , config .max_seq_len , mode = mode , skip_samples = skip_samples )
102+
103+ # Increase num_workers for better utilization
84104 return DataLoader (stream_ds , batch_size = batch_size , pin_memory = True ,
85- num_workers = 1 , prefetch_factor = 2 )
105+ num_workers = 4 , prefetch_factor = 4 )
0 commit comments