(self, batch)
| 9 | """data batch.""" |
| 10 | |
| 11 | def __call__(self, batch): |
| 12 | data_dict = defaultdict(list) |
| 13 | to_tensor_keys = [] |
| 14 | for sample in batch: |
| 15 | for k, v in sample.items(): |
| 16 | if isinstance(v, (np.ndarray, torch.Tensor, numbers.Number)): |
| 17 | if k not in to_tensor_keys: |
| 18 | to_tensor_keys.append(k) |
| 19 | data_dict[k].append(v) |
| 20 | for k in to_tensor_keys: |
| 21 | data_dict[k] = torch.from_numpy(data_dict[k]) |
| 22 | return data_dict |
| 23 | |
| 24 | |
| 25 | class ListCollator(object): |
nothing calls this directly
no outgoing calls
no test coverage detected