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

Method setup

t-few/src/data/data_module.py:21–43  ·  view source on GitHub ↗
(self, stage)

Source from the content-addressed store, hash-verified

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(

Callers

nothing calls this directly

Calls 6

is_custom_taskFunction · 0.90
read_few_shot_datasetMethod · 0.80
read_orig_datasetMethod · 0.45
get_train_templateMethod · 0.45
get_eval_templateMethod · 0.45

Tested by

no test coverage detected