MCPcopy Create free account
hub / github.com/clinicalml/TabLLM / PretrainDataModule

Class PretrainDataModule

t-few/src/data/data_module.py:126–152  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

124
125
126class PretrainDataModule(LightningDataModule):
127 def __init__(self, config, tokenizer, dataset_reader):
128 super().__init__()
129 self.config = config
130 self.tokenizer = tokenizer
131 self.dataset_reader = dataset_reader
132
133 def setup(self, stage):
134 self.train_datasets = self.dataset_reader.read_orig_dataset("train")
135 self.base_templates = self.dataset_reader.get_template()
136 self.train_datasets_withtemplate = []
137 for index, train_dataset in enumerate(self.train_datasets):
138 self.train_datasets_withtemplate.append(
139 PretrainDatasetWithTemplate(train_dataset, self.base_templates[index], self.tokenizer)
140 )
141
142 self.train_dataset = torch.utils.data.ConcatDataset(self.train_datasets_withtemplate)
143
144 def train_dataloader(self):
145 return torch.utils.data.DataLoader(
146 self.train_dataset,
147 batch_size=self.config.batch_size,
148 shuffle=True,
149 collate_fn=create_collate_fn(self.tokenizer.pad_token_id, pretrain=True),
150 drop_last=True,
151 num_workers=min([self.config.batch_size, self.config.num_workers]),
152 )
153
154
155class PretrainDatasetWithTemplate(torch.utils.data.dataset.Dataset):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected