RuntimeError: output with shape x doesn't match the broadcast shape y

Viewed 54

I am trying to get a simple GNN working before I add more functionality to it, but I keep running into the same error:

RuntimeError: output with shape [5969294, 1, 4] doesn't match the broadcast shape [5969294, 5969294, 4] My code is as follows:

import torch
import torch.nn.functional as F
from torch_geometric.nn import TransformerConv, Linear
from torch_geometric.nn import global_mean_pool as gap, global_max_pool as gmp

class GNN(torch.nn.Module):
        def __init__(self, num_node_features, embedding_size):
                super(GNN, self).__init__()

                """ GCN layers"""
                self.conv1 = TransformerConv(num_node_features, embedding_size)
                self.conv2 = TransformerConv(embedding_size, embedding_size)

                """ Output layer """
                self.out = Linear(2*embedding_size, 1)

        def forward(self, x, edge_attr, edge_index, batch_index):
                """ first conv layer """
                x = self.conv1(x, edge_index, edge_attr)
                x = F.relu(x)

                """ second conv layer """
                x = self.conv2(x, edge_index, edge_attr)
                x = F.relu(x)

                """ global pooling using gap and gmp """
                x = torch.cat([gmp(x, batch_index),
                            gap(x, batch_index)], dim=1)

                """ Linear classifier """
                x = self.out(x)

                return x

Then, my training code looks like:

import torch
from model import GNN
from torch_geometric.loader import DataLoader
from process_data import get_data

""" set up device, data, DataLoader, model, loss_fn, and optimizer. Set neccessary variables to Cuda """
device = torch.device('cuda:1' if torch.cuda.is_available() else "cpu")
training = get_data()["training"] # returns a list of torch_geometric Data objects
training_loader = DataLoader(training, batch_size=512, shuffle=True) # high batch size for testing
model = GNN(num_node_features=training[0].x.shape[1], embedding_size=4) # low embedding_size for testing
model = model.to(device)
loss_fn = torch.nn.L1Loss() 
optimizer = torch.optim.SGD(model.parameters(), lr=.01) # high learning rate for testing

def train_epoch(epoch):
    running_loss = 0.0
    step = 0
    for batch in training_loader:
        batch.to(device)
        optimizer.zero_grad()

        """ get pred and its parameters """
        pred_x = batch.x.to(device, dtype=torch.float)
        pred_edge_attr = batch.edge_attr.to(device, dtype=torch.float)
        pred_edge_index = batch.edge_index.to(device, dtype=torch.long)
        pred_batch_index = batch.batch.to(device, dtype=torch.long)
        pred = model(pred_x, pred_edge_attr, pred_edge_index, pred_batch_index) # error arising here

        """ get target """
        target = batch.y.to(device, dtype=torch.float)

        loss = loss_fn(torch.squeeze(pred), torch.squeeze(target))
        loss.backwards()
        optimizer.step()
        running_loss += loss.item()
        step += 1
    return loss/step

""" start training """
best_loss = 10000
for epoch in range(50):
model.train()
    loss = train_epoch(epoch)
    print(f"Epoch {epoch} | Train Loss {loss}")
    if float(loss) < best_loss:
        best_loss = loss
print(f"Best loss was {best_loss}")

Here is the full output of the error:

Traceback (most recent call last): File "/home/cfalkenberg/train2.py", line 43, in loss = train_epoch(epoch)
File "/home/cfalkenberg/train2.py", line 27, in train_epoch
pred = model(pred_x, pred_edge_attr, pred_edge_index, pred_batch_index)
File "/home/cfalkenberg/anaconda3/envs/cfalk/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1110, in _call_impl
return forward_call(*input, **kwargs).
File "/home/cfalkenberg/model.py", line 24, in forward
x = self.conv1(x, edge_index, edge_attr)
File "/home/cfalkenberg/anaconda3/envs/cfalk/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1110, in _call_impl
return forward_call(*input, **kwargs)
File "/home/cfalkenberg/anaconda3/envs/cfalk/lib/python3.9/site-packages/torch_geometric/nn/conv/transformer_conv.py", line 176, in forward
out = self.propagate(edge_index, query=query, key=key, value=value,
File "/home/cfalkenberg/anaconda3/envs/cfalk/lib/python3.9/site-packages/torch_geometric/nn/conv/message_passing.py", line 317, in propagate
out = self.message(**msg_kwargs)
File "/home/cfalkenberg/anaconda3/envs/cfalk/lib/python3.9/site-packages/torch_geometric/nn/conv/transformer_conv.py", line 222, in message
out += edge_attr
RuntimeError: output with shape [5969294, 1, 4] doesn't match the broadcast shape [5969294, 5969294, 4]

I think there is somewhere that my input shape is incorrect, but I am really confused as to where. I believe the problem arises at the line where I define the "pred" variables

Does anyone know what the issue is?

0 Answers
Related