(self, tokenizer, data_args, tpu_num_cores=None)
| 274 | |
| 275 | class Seq2SeqDataCollator: |
| 276 | def __init__(self, tokenizer, data_args, tpu_num_cores=None): |
| 277 | self.tokenizer = tokenizer |
| 278 | self.pad_token_id = tokenizer.pad_token_id |
| 279 | assert ( |
| 280 | self.pad_token_id is not None |
| 281 | ), f"pad_token_id is not defined for ({self.tokenizer.__class__.__name__}), it must be defined." |
| 282 | self.data_args = data_args |
| 283 | self.tpu_num_cores = tpu_num_cores |
| 284 | self.dataset_kwargs = {"add_prefix_space": True} if isinstance(tokenizer, BartTokenizer) else {} |
| 285 | if data_args.src_lang is not None: |
| 286 | self.dataset_kwargs["src_lang"] = data_args.src_lang |
| 287 | if data_args.tgt_lang is not None: |
| 288 | self.dataset_kwargs["tgt_lang"] = data_args.tgt_lang |
| 289 | |
| 290 | def __call__(self, batch) -> Dict[str, torch.Tensor]: |
| 291 | if hasattr(self.tokenizer, "prepare_seq2seq_batch"): |
nothing calls this directly
no outgoing calls
no test coverage detected