1010
1111from fast_llm .data .blended import BlendedDataset
1212from fast_llm .data .config import Data , DatasetSource , SampledDataset
13- from fast_llm .data .gpt .config import DataConfig
13+ from fast_llm .data .gpt .config import GPTDataConfig
14+ from fast_llm .data .gpt .dataset import GPTSamplingConfig
1415from fast_llm .data .gpt .dummy import DummyGPTDataset
1516from fast_llm .data .gpt .memmap import GPTMemmapDataset
16- from fast_llm .data .gpt .sampled import GPTSampledIndexedDataset
1717from fast_llm .data .gpt .slice import GPTDatasetSlice
1818from fast_llm .data .iterator import SampledDatasetIterator
1919from fast_llm .data .tokenizer import Tokenizer
@@ -40,10 +40,11 @@ class GPTData(Data):
4040 _cache_directory : pathlib .Path | None
4141 _samples_per_phase : dict [PhaseType , int ]
4242 _phases : typing .ClassVar [tuple [PhaseType , ...]] = (PhaseType .training , PhaseType .validation , PhaseType .test )
43+ _is_setup : bool = False
4344
4445 def __init__ (
4546 self ,
46- config : DataConfig ,
47+ config : GPTDataConfig ,
4748 distributed_config : DistributedConfig ,
4849 vocab_size : int ,
4950 max_sequence_length : int ,
@@ -52,8 +53,8 @@ def __init__(
5253 Create the data and gather some basic information on the dataset(s).
5354 Should be `setup` before use.
5455 """
55- self ._config = config . validate ()
56- self ._distributed_config = distributed_config . validate ()
56+ self ._config = config
57+ self ._distributed_config = distributed_config
5758 self ._vocab_size = vocab_size
5859 self ._max_sequence_length = max_sequence_length
5960 Assert .eq (len (self ._config .split ), len (self ._phases ))
@@ -166,6 +167,20 @@ def setup(self, distributed: Distributed, samples_per_phase: dict[PhaseType, int
166167 )
167168 for phase , datasets in self ._sampled_datasets .items ()
168169 }
170+ self ._is_setup = True
171+
172+ @property
173+ def config (self ):
174+ return self ._config
175+
176+ @property
177+ def tokenizer (self ):
178+ assert self ._is_setup
179+ return self ._tokenizer
180+
181+ @property
182+ def distributed (self ):
183+ return self ._distributed
169184
170185 def get_iterator (
171186 self ,
@@ -176,6 +191,7 @@ def get_iterator(
176191 num_workers : int ,
177192 prefetch_factor : int | None = None ,
178193 ):
194+ assert self ._is_setup
179195 Assert .incl (phase , self ._blended_datasets )
180196 Assert .in_range_incl (batch_config .sequence_length , 1 , self ._max_sequence_length )
181197 log_main_rank (f"Initializing { phase } data iterator from sample { consumed_samples } ..." )
@@ -205,25 +221,23 @@ def _build_and_sample_gpt_dataset(self, name: str, dataset_samples_per_phase: di
205221 for phase , num_samples in dataset_samples_per_phase .items ():
206222 if num_samples == 0 :
207223 continue
208- sampled_datasets [phase ] = GPTSampledIndexedDataset (
209- dataset_split [phase ],
210- num_samples = num_samples ,
211- sequence_length = self ._max_sequence_length ,
212- seed = self ._distributed .config .seed ,
213- group = self ._distributed .world_group ,
214- config = self ._config ,
215- tokenizer = self ._tokenizer ,
216- cache_directory = (
217- self ._dataset_prefixes [name ].parent if self ._cache_directory is None else self ._cache_directory
224+ sampled_datasets [phase ] = dataset_split [phase ].sample (
225+ GPTSamplingConfig (
226+ num_samples = num_samples ,
227+ sequence_length = self ._max_sequence_length ,
228+ seed = self ._distributed_config .seed ,
229+ cache_directory = (
230+ self ._dataset_prefixes [name ].parent if self ._cache_directory is None else self ._cache_directory
231+ ),
232+ verbose = self ._num_datasets <= 5 ,
218233 ),
219- verbose = self . _num_datasets <= 5 ,
234+ self ,
220235 )
221236 return sampled_datasets
222237
223238 def _build_and_sample_dummy_dataset (self , name : str , dataset_samples_per_phase : dict [PhaseType , int ]):
224239 return {
225240 phase : DummyGPTDataset (
226- self ._dataset_prefixes [name ],
227241 dataset_samples_per_phase [phase ],
228242 self ._max_sequence_length ,
229243 self ._vocab_size ,
0 commit comments