(cls: type[Self], data: dict[T_Fields, torch.Tensor])
| 31 | |
| 32 | @classmethod |
| 33 | def init(cls: type[Self], data: dict[T_Fields, torch.Tensor]) -> Self: |
| 34 | sizes = [v.size(0) for v in data.values()] |
| 35 | assert all([s == sizes[0] for s in sizes]), f"TensorBundle requires all features have (Nx...) shape with same 'N', get {sizes}" |
| 36 | |
| 37 | index = torch.arange(0, sizes[0], dtype=torch.long) |
| 38 | return cls(index, data) |
| 39 | |
| 40 | def __getitem__(self, index) -> TensorBundle[T_Fields]: |
| 41 | selected_dict: dict[T_Fields, torch.Tensor] = { |
no outgoing calls
no test coverage detected