| 124 | |
| 125 | |
| 126 | class 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 | |
| 155 | class PretrainDatasetWithTemplate(torch.utils.data.dataset.Dataset): |
nothing calls this directly
no outgoing calls
no test coverage detected