| 120 | |
| 121 | # -------------------------------------------- Graph specific transforms |
| 122 | class 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 | |
| 174 | class OnlyLayersAreNodesTransforms(AbstractGraphTransform, ABC): |
nothing calls this directly
no outgoing calls
no test coverage detected