(self, batch_samples: List[Dict[str, Any]])
| 335 | return output |
| 336 | |
| 337 | def batchify(self, batch_samples: List[Dict[str, Any]]) -> Dict[str, Any]: |
| 338 | batch = dict() |
| 339 | batch_text_vec = torch.tensor(pad_sequences( |
| 340 | [sample['text_vec'] for sample in batch_samples], pad_value=self.tokenizer.pad_token_id, pad_left=False |
| 341 | ), dtype=torch.long) |
| 342 | loss_mask = torch.tensor(pad_sequences( |
| 343 | [sample['loss_mask'] for sample in batch_samples], pad_value=0, pad_left=False |
| 344 | ), dtype=torch.bool) |
| 345 | |
| 346 | batch.update({ |
| 347 | 'text_vec': batch_text_vec, |
| 348 | 'loss_mask': loss_mask |
| 349 | }) |
| 350 | |
| 351 | return batch |
| 352 | |
| 353 | def batch_generator(self): |
| 354 | while True: |
no test coverage detected