MCPcopy Create free account
hub / github.com/MAC-VO/MAC-VO / collate

Method collate

DataLoader/Interface.py:19–48  ·  view source on GitHub ↗

A default collate function that will handle torch.Tensor, pp.LieTensor and np.array automatically. You can perform more customized collate by one of the following methods: 1. Overriding the collate method 2. Setting the class attribute `collate_hand

(cls, batch: T.Sequence[Self])

Source from the content-addressed store, hash-verified

17
18 @classmethod
19 def collate(cls, batch: T.Sequence[Self]) -> Self:
20 """
21 A default collate function that will handle torch.Tensor, pp.LieTensor and
22 np.array automatically. You can perform more customized collate by one of the following methods:
23
24 1. Overriding the collate method
25
26 2. Setting the class attribute `collate_handlers` to a dictionary that maps the attribute name to the collate function corresponding to that field.
27
28 """
29 data_dict = dict()
30 for key, value in batch[0].__dict__.items():
31 if key in cls.collate_handlers:
32 collate_fn = cls.collate_handlers[key]
33 elif isinstance(value, torch.Tensor):
34 collate_fn = lambda seq: torch.cat(seq, dim=0)
35 elif isinstance(value, pp.LieTensor):
36 collate_fn = lambda seq: torch.stack(seq, dim=0)
37 elif isinstance(value, np.ndarray):
38 collate_fn = lambda seq: np.concatenate(seq, axis=0)
39 elif isinstance(value, list):
40 collate_fn = lambda seq: list(chain.from_iterable(seq))
41 elif isinstance(value, Collatable):
42 collate_fn = value.collate
43 elif value is None:
44 collate_fn = lambda seq: None
45 else:
46 raise ValueError(f"Unsupported data type {type(value)}, you need to overrider the collate method.")
47 data_dict[key] = cls._collate([getattr(x, key) for x in batch], collate_fn)
48 return cls(**data_dict)
49
50 @staticmethod
51 def _collate(batch: T.Sequence[Tp | None], collate_fn: CollateFn) -> Tp | None:

Callers 1

collate_fnMethod · 0.80

Calls 1

_collateMethod · 0.80

Tested by

no test coverage detected