MCPcopy Create free account
hub / github.com/Ugness/ELF-pytorch / ELFDataModule

Class ELFDataModule

pytorch_lightning/lightning_module.py:352–378  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

350# DataModule
351# -----------------------------------------------------------------------------
352class 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 )

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected