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

Class AbstractGraphTransform

climart/data_transform/transforms.py:122–171  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

120
121# -------------------------------------------- Graph specific transforms
122class AbstractGraphTransform(AbstractTransform, ABC):
123 def __init__(self, exp_type: str):
124 super().__init__(exp_type)
125
126 @property
127 def n_nodes(self) -> int:
128 return self._n_nodes
129
130 def get_edge_idxs(self) -> Tuple[np.ndarray, np.ndarray]:
131 # one-way: node i has an edge to node i+1
132 one_way_senders = np.arange(self.n_nodes - 1) # n_lay - 1 = n_lev - 2 edges
133 one_way_receivers = one_way_senders + 1
134 # one-way: node i+1 has an edge to node i
135 other_way_senders = np.arange(1, self.n_nodes)
136 other_way_receivers = other_way_senders - 1
137
138 senders = np.concatenate((one_way_senders, other_way_senders))
139 receivers = np.concatenate((one_way_receivers, other_way_receivers))
140 return senders, receivers
141
142 def get_adj(self, degree_normalized: bool = False, improved: bool = False) -> Tensor:
143 """ Adjacency matrix of a line graph with global node and self-loops """
144 adj = torch.zeros((self.n_nodes, self.n_nodes))
145
146 if hasattr(self, 'GLOBAL_NODE'):
147 adj[:, self.GLOBAL_NODE] = 1
148 adj[self.GLOBAL_NODE, :] = 1
149
150 for i in range(1, self.n_nodes):
151 adj[i, i - 1:i + 2] = 1
152 adj[i - 1:i + 2, i] = 1
153
154 if degree_normalized:
155 self.log.info("-> Adjacency matrix is normalized by in-degree")
156 return normalize_adjacency_matrix_torch(adj, improved=improved, add_self_loops=True)
157 return adj
158
159 def get_level_nodes_mask(self) -> Tensor:
160 """
161 Indexing an array 'A' with shape (n, self.n_nodes, out-dim)
162 leads to A[:, get_level_nodes_mask(), :] have shape (n, len(self.LEVEL_NODES), out-dim)
163 :return: A torch.Tensor mask that indexes all the level nodes
164 """
165
166 if not hasattr(self, 'LEVEL_NODES'):
167 raise ValueError(f"This transform {self} does not have level nodes.")
168 else:
169 level_mask = torch.zeros(self.n_nodes)
170 level_mask[self.LEVEL_NODES] = 1
171 return level_mask.bool()
172
173
174class OnlyLayersAreNodesTransforms(AbstractGraphTransform, ABC):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected