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

Class DataCollatorForSupervisedDataset

train/train.py:215–252  ·  view source on GitHub ↗

Collate examples for supervised fine-tuning.

Source from the content-addressed store, hash-verified

213
214@dataclass
215class DataCollatorForSupervisedDataset(object):
216 """Collate examples for supervised fine-tuning."""
217
218 tokenizer: transformers.PreTrainedTokenizer
219
220 def naive__call__(self, instances: Sequence[Dict]) -> Dict[str, torch.Tensor]:
221 input_ids, labels = tuple([instance[key] for instance in instances] for key in ("input_ids", "labels"))
222 input_ids = torch.nn.utils.rnn.pad_sequence(
223 input_ids, batch_first=True, padding_value=self.tokenizer.pad_token_id
224 )
225 labels = torch.nn.utils.rnn.pad_sequence(labels, batch_first=True, padding_value=IGNORE_INDEX)
226 return dict(
227 input_ids=input_ids,
228 labels=labels,
229 attention_mask=input_ids.ne(self.tokenizer.pad_token_id),
230 )
231
232 def __call__(self, instances: Sequence[Dict]) -> Dict[str, torch.Tensor]:
233 sources = []
234 targets = []
235 for instance in instances:
236 source = instance['input_ids']
237 target = instance['labels']
238 sources.append(source)
239 targets.append(target)
240
241 data_dict = preprocess(sources, targets, self.tokenizer)
242 input_ids, labels = data_dict['input_ids'], data_dict['labels']
243
244 input_ids = torch.nn.utils.rnn.pad_sequence(
245 input_ids, batch_first=True, padding_value=self.tokenizer.pad_token_id
246 )
247 labels = torch.nn.utils.rnn.pad_sequence(labels, batch_first=True, padding_value=IGNORE_INDEX)
248 return dict(
249 input_ids=input_ids,
250 labels=labels,
251 attention_mask=input_ids.ne(self.tokenizer.pad_token_id),
252 )
253
254def make_supervised_data_module(tokenizer: transformers.PreTrainedTokenizer, data_args) -> Dict:
255 """Make dataset and collator for supervised fine-tuning."""

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected