I am currently implementing a variant of GAT (Graph Attention Network) which also takes into account edges features, trying to reproduce this paper https://arxiv.org/abs/2101.07671.
I took as inspiration this GAT implementation https://nn.labml.ai/graphs/gat/index.html, with the only difference with the only difference in the received inputs, which include the features of the arcs, and the adjacency matrices not represented as a sparse NxN matrix but a compressed version which only keeps track of connected nodes.
The attention score is computed taking into account feeatures concatenatio of nodes and edges and since I am working with big graphs, this leads to the creation of too large matrices, which cannot be loaded on the GPU.
This is my code:
linear_transformed_nodes_repeated = linear_transformed_nodes.repeat(n_nodes, 1)
linear_transformed_nodes_repeated_interleave = linear_transformed_nodes.repeat_interleave(n_nodes, dim=0)
#Node concatenation
linear_transformed_nodes_concat = torch.cat([linear_transformed_nodes_repeated_interleave, linear_transformed_nodes_repeated], dim=-1)
#Each concatenation is now repeated n_edges times becuase will be concatenated with every edge
inear_transformed_nodes_concat = linear_transformed_nodes_concat.repeat_interleave(n_edges, dim = 0)
linear_transformed_edges_repeated = linear_transformed_edges.repeat(n_nodes * n_nodes,1)
#Node-Node-Edge concatenation
nodes_edge_concatenation = torch.cat([linear_transformed_nodes_concat, linear_transformed_edges_repeated], dim=1)
#Reshape the matrix so that A[x][y][z] contains the concatenation between nodes x, y and edge z
nodes_edge_concatenation = nodes_edge_concatenation.view(n_nodes, n_nodes, n_edges, self.out_features_nodes * 2 + self.out_features_edges)
e = self.activation(self.attn_nodes(nodes_edge_concatenation))
a = self.softmax(e)
self.attn_nodes is just a linear transformation, while self.activation is a LeakyReLU.
After that, i simply produce the new node embedding by summing over its neighborhood.
The main problem of this implementation is the size of the concatenation matrices, for example a graph with 800 nodes and 66 edges produces after the 4th transformation a matrix with shape [42240000, 400], assuming each 200 features per node. Is there a way to compute the attention score without computing all these intermediate tensors?