(self, data_list)
| 290 | return out |
| 291 | |
| 292 | def _batch(self, data_list): |
| 293 | batch_data = {} |
| 294 | for sample_preprocessed in data_list: |
| 295 | for k, v in sample_preprocessed.items(): |
| 296 | value_list = batch_data.get(k, []) |
| 297 | value_list.append(v) |
| 298 | batch_data[k] = value_list |
| 299 | for k in batch_data.keys(): |
| 300 | if isinstance(batch_data[k][0], torch.Tensor): |
| 301 | batch_data[k] = torch.cat(batch_data[k]) |
| 302 | return batch_data |
| 303 | |
| 304 | def _process_batch(self, input: List[Input], batch_size, |
| 305 | **kwargs) -> Dict[str, Any]: |
no test coverage detected