Collate examples for supervised fine-tuning.
| 213 | |
| 214 | @dataclass |
| 215 | class 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 | |
| 254 | def make_supervised_data_module(tokenizer: transformers.PreTrainedTokenizer, data_args) -> Dict: |
| 255 | """Make dataset and collator for supervised fine-tuning.""" |
no outgoing calls
no test coverage detected