(self, batch)
| 26 | """data batch.""" |
| 27 | |
| 28 | def __call__(self, batch): |
| 29 | data_dict = defaultdict(list) |
| 30 | to_tensor_idxs = [] |
| 31 | for sample in batch: |
| 32 | for idx, v in enumerate(sample): |
| 33 | if isinstance(v, (np.ndarray, torch.Tensor, numbers.Number)): |
| 34 | if idx not in to_tensor_idxs: |
| 35 | to_tensor_idxs.append(idx) |
| 36 | data_dict[idx].append(v) |
| 37 | for idx in to_tensor_idxs: |
| 38 | data_dict[idx] = torch.from_numpy(data_dict[idx]) |
| 39 | return list(data_dict.values()) |
| 40 | |
| 41 | |
| 42 | class SSLRotateCollate(object): |
nothing calls this directly
no outgoing calls
no test coverage detected