Obtain cross edge index that fully connect nodes of corresponding graphs in two different DataBatch

Viewed 30

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]])
0 Answers
Related