Creates a TensorStatistics object from a tensor.
(cls, tensor: torch.Tensor)
| 28 | |
| 29 | @classmethod |
| 30 | def from_tensor(cls, tensor: torch.Tensor) -> "TensorStatistics": |
| 31 | """Creates a TensorStatistics object from a tensor.""" |
| 32 | flattened = torch.flatten(tensor) |
| 33 | return cls( |
| 34 | shape=tensor.shape, |
| 35 | numel=tensor.numel(), |
| 36 | median=torch.median(flattened).item(), |
| 37 | mean=flattened.mean().item(), |
| 38 | max=flattened.max().item(), |
| 39 | min=flattened.min().item(), |
| 40 | ) |
| 41 | |
| 42 | |
| 43 | @dataclass |