MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / SupervisedDataset

Class SupervisedDataset

data/generation/single_generate.py:93–109  ·  view source on GitHub ↗

Dataset for supervised fine-tuning.

Source from the content-addressed store, hash-verified

91 return dict(input_ids=input_ids, labels=copy.deepcopy(input_ids))
92
93class SupervisedDataset(Dataset):
94 """Dataset for supervised fine-tuning."""
95 def __init__(self, dataset_name: str, tokenizer: transformers.PreTrainedTokenizer, max_sample=None):
96 super(SupervisedDataset, self).__init__()
97
98 sources, targets = get_gen_dataset(dataset_name, max_sample, tokenizer)
99
100 data_dict = preprocess(sources, targets, tokenizer)
101
102 self.input_ids = data_dict["input_ids"]
103 self.labels = data_dict["labels"]
104
105 def __len__(self):
106 return len(self.input_ids)
107
108 def __getitem__(self, i) -> Dict[str, torch.Tensor]:
109 return dict(input_ids=self.input_ids[i], labels=self.labels[i], id=i)
110
111def padding(inputs, padding_token, cutoff = None):
112 num_elems = len(inputs)

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected