Collate examples for supervised fine-tuning.
| 124 | |
| 125 | @dataclass |
| 126 | class DataCollatorForSupervisedDataset(object): |
| 127 | """Collate examples for supervised fine-tuning.""" |
| 128 | |
| 129 | tokenizer: transformers.PreTrainedTokenizer |
| 130 | |
| 131 | def __call__(self, instances: Sequence[Dict]) -> Dict[str, torch.Tensor]: |
| 132 | input_ids, labels, ids = tuple([instance[key] for instance in instances] for key in ("input_ids", "labels", 'id')) |
| 133 | input_ids = padding(input_ids, self.tokenizer.pad_token_id, cutoff = 256) |
| 134 | labels = padding(labels, IGNORE_INDEX, cutoff = 256) |
| 135 | |
| 136 | return dict( |
| 137 | input_ids=input_ids, |
| 138 | labels=labels, |
| 139 | id=torch.tensor(ids).to(input_ids.device), |
| 140 | attention_mask=input_ids.ne(self.tokenizer.pad_token_id), |
| 141 | ) |
| 142 | |
| 143 | def make_supervised_data_module(tokenizer: transformers.PreTrainedTokenizer, dataset_name, max_sample=None) -> Dict: |
| 144 | """Make dataset and collator for supervised fine-tuning.""" |
no outgoing calls
no test coverage detected