I have two DataBatch for which I would like to fully connect the graphs, in a pairwise manner. By that I mean that if I have DataBatch G (made of the graphs g1,g2,g3,...) and H (h1, h2, h3, ...), I want to get a cross_edge_index that links all nodes from g1 to all nodes from h1, all nodes from g2 to all nodes from h2, etc.
I can do it by iteratively looping through each batch, but for efficiency, I would like to use only tensor operations.
Toy example:
from torch_geometric.data.batch import Batch
from torch_geometric.data import Data
g1 = Data(x=torch.ones((5))) # Graph with 5 nodes
g2 = Data(x=torch.ones((4))) # Graph with 4 nodes
G = Batch.from_data_list([g1,g2])
h1 = Data(x=torch.ones((2))) # Graph with 2 nodes
h2 = Data(x=torch.ones((3))) # Graph with 3 nodes
H = Batch.from_data_list([h1,h2])
Expected output:
cross_edge_index =
tensor([[0, 0, 1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 5, 6, 6, 6, 7, 7, 7, 8, 8, 8],
[0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 2, 3, 4, 2, 3, 4, 2, 3, 4, 2, 3, 4]])
My (ugly) working solution (using for-loop):
import torch
def get_cross_edge_index(g, h):
edge_index = []
for g_, h_, offset_g, offset_h in zip(g.to_data_list(), h.to_data_list(), g.ptr[:-1], h.ptr[:-1]):
edge_index.append(torch.cartesian_prod(offset_g + torch.arange(g_.num_nodes), offset_h + torch.arange(h_.num_nodes)).T)
return torch.cat(edge_index, axis=1)
get_cross_edge_index(G,H)
>>> tensor([[0, 0, 1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 5, 6, 6, 6, 7, 7, 7, 8, 8, 8],
[0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 2, 3, 4, 2, 3, 4, 2, 3, 4, 2, 3, 4]])