(self, stage)
| 19 | _ = self.dataset_reader.read_few_shot_dataset() |
| 20 | |
| 21 | def setup(self, stage): |
| 22 | # make assignments here (val/train/test split) |
| 23 | # called on every process in DDP |
| 24 | if self.config.few_shot: |
| 25 | self.train_dataset = self.dataset_reader.read_few_shot_dataset() |
| 26 | else: |
| 27 | self.train_dataset = self.dataset_reader.read_orig_dataset("train") |
| 28 | self.dev_dataset = self.dataset_reader.read_orig_dataset("validation") |
| 29 | self.train_dataset = FinetuneDatasetWithTemplate( |
| 30 | self.train_dataset, self.dataset_reader.get_train_template(), self.tokenizer |
| 31 | ) |
| 32 | self.dev_dataset = FinetuneDatasetWithTemplate( |
| 33 | self.dev_dataset, self.dataset_reader.get_eval_template(), self.tokenizer |
| 34 | ) |
| 35 | print(f"Train size {len(self.train_dataset)}") |
| 36 | print(f"Eval size {len(self.dev_dataset)}") |
| 37 | |
| 38 | if is_custom_task(self.config): |
| 39 | self.test_dataset = self.dataset_reader.read_orig_dataset("test") |
| 40 | self.test_dataset = FinetuneDatasetWithTemplate( |
| 41 | self.test_dataset, self.dataset_reader.get_eval_template(), self.tokenizer |
| 42 | ) |
| 43 | print(f"Test size {len(self.test_dataset)}") |
| 44 | |
| 45 | def train_dataloader(self): |
| 46 | return torch.utils.data.DataLoader( |
nothing calls this directly
no test coverage detected