graph_net_level_nodes: gn_input_dict_renamer_level_nodes
| 309 | |
| 310 | |
| 311 | class 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 | |
| 334 | class LevelEdgesAndLayerNodesGraph(EdgesAndNodesTransforms, OnlyLayersAreNodesTransforms): |
nothing calls this directly
no outgoing calls
no test coverage detected