| 39 | |
| 40 | |
| 41 | class PredictionPostProcessCallback: |
| 42 | def __init__(self, |
| 43 | variables: List[str], |
| 44 | sizes: Union[int, Sequence[int]] |
| 45 | ): |
| 46 | self.variable_to_channel = dict() |
| 47 | cur = 0 |
| 48 | sizes = [sizes for _ in range(len(variables))] if isinstance(sizes, int) else sizes |
| 49 | for var, size in zip(variables, sizes): |
| 50 | self.variable_to_channel[var] = {'start': cur, 'end': cur + size} |
| 51 | cur += size |
| 52 | |
| 53 | def split_vector_by_variable(self, |
| 54 | vector: Union[np.ndarray, torch.Tensor] |
| 55 | ) -> Dict[str, Union[np.ndarray, torch.Tensor]]: |
| 56 | if isinstance(vector, dict): |
| 57 | return vector |
| 58 | splitted_vector = dict() |
| 59 | for var_name, var_channel_limits in self.variable_to_channel.items(): |
| 60 | splitted_vector[var_name] = vector[..., var_channel_limits['start']:var_channel_limits['end']] |
| 61 | return splitted_vector |
| 62 | |
| 63 | def __call__(self, vector, *args, **kwargs): |
| 64 | return self.split_vector_by_variable(vector) |