Collate examples for supervised fine-tuning.
| 151 | |
| 152 | @dataclass |
| 153 | class DataCollatorForSupervisedDataset(object): |
| 154 | """Collate examples for supervised fine-tuning.""" |
| 155 | |
| 156 | tokenizer: transformers.PreTrainedTokenizer |
| 157 | |
| 158 | def __call__(self, instances: Sequence[Dict]) -> Dict[str, torch.Tensor]: |
| 159 | input_ids, labels = tuple([instance[key] for instance in instances] for key in ("input_ids", "labels")) |
| 160 | input_ids = [torch.tensor(x) for x in input_ids] |
| 161 | input_ids = torch.nn.utils.rnn.pad_sequence( |
| 162 | input_ids, batch_first=True, padding_value=self.tokenizer.pad_token_id |
| 163 | ) |
| 164 | labels = [torch.tensor(x) for x in labels] |
| 165 | labels = torch.nn.utils.rnn.pad_sequence(labels, batch_first=True, padding_value=IGNORE_INDEX) |
| 166 | return dict( |
| 167 | input_ids=input_ids, |
| 168 | labels=labels, |
| 169 | attention_mask=input_ids.ne(self.tokenizer.pad_token_id), |
| 170 | ) |
| 171 | |
| 172 | def _tokenize_fn(strings: Sequence[str], tokenizer: transformers.PreTrainedTokenizer) -> Dict: |
| 173 | """Tokenize a list of strings.""" |