An arbitrary one-to-multi mapping relationship.
| 131 | |
| 132 | |
| 133 | class SparseEdge_Multi(EdgeLike): |
| 134 | """ |
| 135 | An arbitrary one-to-multi mapping relationship. |
| 136 | """ |
| 137 | def __init__(self, num_from: int, max_deg: int): |
| 138 | self.out_deg = torch.zeros((num_from,), dtype=torch.long) |
| 139 | self.edges = torch.ones((num_from, max_deg), dtype=torch.long) * -1 |
| 140 | self.max_deg = max_deg |
| 141 | |
| 142 | def add(self, from_idx: torch.Tensor, to_idx: torch.Tensor): |
| 143 | assert from_idx.shape == to_idx.shape |
| 144 | self.edges[from_idx, self.out_deg[from_idx]] = to_idx |
| 145 | self.out_deg[from_idx] += 1 |
| 146 | |
| 147 | def project(self, from_index: torch.Tensor) -> torch.Tensor: |
| 148 | to_idx = self.edges[from_index].flatten() |
| 149 | return to_idx[to_idx >= 0] |
| 150 | |
| 151 | def serialize(self, prefix: str) -> dict[str, np.ndarray]: |
| 152 | return { |
| 153 | f"{prefix}/edges": self.edges.cpu().numpy(), |
| 154 | f"{prefix}/deg" : self.out_deg.cpu().numpy() |
| 155 | } |
| 156 | |
| 157 | @classmethod |
| 158 | def deserialize(cls, prefix: str, value: dict[str, np.ndarray]) -> Self: |
| 159 | edges = torch.Tensor(value[f"{prefix}/edges"]) |
| 160 | deg = torch.Tensor(value[f"{prefix}/deg"]) |
| 161 | num_from, max_deg = edges.shape[0], edges.shape[1] |
| 162 | |
| 163 | edge_instance = cls(num_from, max_deg) |
| 164 | edge_instance.edges = edges |
| 165 | edge_instance.out_deg = deg |
| 166 | return edge_instance |
| 167 | |
| 168 | |
| 169 | class DenseEdge_Multi(EdgeLike): |