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])
| 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: |