MCPcopy Create free account
hub / github.com/RolnickLab/climart / PredictionPostProcessCallback

Class PredictionPostProcessCallback

climart/utils/callbacks.py:41–64  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

39
40
41class 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)

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected