| 350 | # DataModule |
| 351 | # ----------------------------------------------------------------------------- |
| 352 | class ELFDataModule(L.LightningDataModule): |
| 353 | def __init__(self, config: Config, tokenizer): |
| 354 | super().__init__() |
| 355 | self.cfg = config |
| 356 | self.tokenizer = tokenizer |
| 357 | self._train_dataset = None |
| 358 | self._eval_dataset = None |
| 359 | |
| 360 | def setup(self, stage=None): |
| 361 | if self._train_dataset is None: |
| 362 | self._train_dataset, self._eval_dataset = load_dataset(self.cfg) |
| 363 | |
| 364 | def train_dataloader(self): |
| 365 | # Lightning attaches the DistributedSampler when strategy=ddp. |
| 366 | return make_dataloader( |
| 367 | self._train_dataset, |
| 368 | batch_size=self.cfg.global_batch_size // self.trainer.world_size, |
| 369 | shuffle=True, |
| 370 | max_seq_length=self.cfg.max_length, |
| 371 | pad_token_id=get_pad_token_id(self.tokenizer, self.cfg.pad_token), |
| 372 | max_input_seq_length=self.cfg.max_input_length, |
| 373 | num_workers=self.cfg.num_workers, |
| 374 | prefetch_factor=self.cfg.prefetch_factor, |
| 375 | pin_memory=self.cfg.pin_memory, |
| 376 | persistent_workers=self.cfg.persistent_workers, |
| 377 | drop_last=True, |
| 378 | ) |