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

Class LayerEdgesAndLevelNodesGraph

climart/data_transform/transforms.py:311–331  ·  view source on GitHub ↗

graph_net_level_nodes: gn_input_dict_renamer_level_nodes

Source from the content-addressed store, hash-verified

309
310
311class LayerEdgesAndLevelNodesGraph(EdgesAndNodesTransforms, OnlyLevelsAreNodesTransforms):
312 """
313 graph_net_level_nodes: gn_input_dict_renamer_level_nodes
314 """
315
316 def __init__(self, exp_type: str):
317 super().__init__(exp_type)
318 self._n_edges = self.spatial_input_dim[LAYERS] * 2 # bi-directional
319 self._out_dim = {NODES: self.input_dim[LEVELS], EDGES: self.input_dim[LAYERS], GLOBALS: self.input_dim[GLOBALS]}
320
321 def transform(self, X: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
322 X[NODES] = X.pop(LEVELS)
323 X[EDGES] = X.pop(LAYERS)
324 X[EDGES] = einops.repeat(X[EDGES], "e d -> (repeat e) d", repeat=2) # bidirectional edges
325 return X
326
327 def batched_transform(self, X: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
328 X[NODES] = X.pop(LEVELS)
329 X[EDGES] = X.pop(LAYERS) # x[GLOBALS] = x.pop(GLOBALS)
330 X[EDGES] = einops.repeat(X[EDGES], "b e d -> b (repeat e) d", repeat=2) # bidirectional edges
331 return X
332
333
334class LevelEdgesAndLayerNodesGraph(EdgesAndNodesTransforms, OnlyLayersAreNodesTransforms):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected