(edge_index, size)
| 44 | |
| 45 | @staticmethod |
| 46 | def norm(edge_index, size): |
| 47 | edge_weight = tf.ones([tf.shape(edge_index)[1], 1]) |
| 48 | |
| 49 | def deg_inv_sqrt(i): |
| 50 | deg = mp_ops.scatter_add(edge_weight, edge_index[i], size[i]) |
| 51 | return deg ** -0.5 |
| 52 | |
| 53 | return tuple(map(deg_inv_sqrt, [0, 1])) |
| 54 | |
| 55 | def __call__(self, x, edge_index, size=None, **kwargs): |
| 56 | norm = self.norm(edge_index, size) |