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

Class SupervisedDataset

train/train.py:175–212  ·  view source on GitHub ↗

Dataset for supervised fine-tuning.

Source from the content-addressed store, hash-verified

173
174
175class SupervisedDataset(Dataset):
176 """Dataset for supervised fine-tuning."""
177
178 def __init__(self, data_path: str, tokenizer: transformers.PreTrainedTokenizer, max_sample: int, split: str):
179 super().__init__()
180
181 with open(data_path, 'r') as f:
182 lines = f.readlines()
183 all_dataset = [json.loads(line.strip()) for line in lines]
184
185 sources, targets = zip(*[(s[0][0], f"{s[0][1]}{tokenizer.eos_token}") for s in all_dataset])
186
187 dataset_size = len(sources)
188 max_sample = min(max_sample or dataset_size, dataset_size)
189 if max_sample < dataset_size:
190 indices = random.sample(range(dataset_size), max_sample)
191 self.sources, self.targets = [sources[i] for i in indices], [targets[i] for i in indices]
192 else:
193 self.sources, self.targets = sources, targets
194
195 split_num = len(self.sources) // 5
196 if split == "train":
197 self.sources, self.targets = self.sources[split_num:], self.targets[split_num:]
198 print(f"Using {len(self.sources)} samples to train")
199
200 print("Example Data")
201 print("sources: \n", self.sources[0])
202 print("targets: \n", self.targets[0])
203
204 elif split == "eval":
205 self.sources, self.targets = self.sources[:split_num], self.targets[:split_num]
206 print(f"Using {len(self.sources)} samples to evaluation")
207
208 def __len__(self):
209 return len(self.sources)
210
211 def __getitem__(self, i):
212 return dict(input_ids=self.sources[i], labels=self.targets[i])
213
214@dataclass
215class DataCollatorForSupervisedDataset(object):

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected