I am trying to do a graph convolution network with forward-forward propagation. When I am running the model, I get an error saying that my index is out of bounds, but this happens after the model works fine for some examples. My dataset contains 253 graphs and the batch size is 64. I don’t understand what would be the problem for one graph but not for the others. Thank you!
import torch
import torch.nn.functional as F
import subprocess
import torch.nn as nn
from torch.nn import MSELoss
from torch_geometric.nn import GATConv, GCNConv
import numpy as np
from torch_geometric.data import Data
from torch_geometric.loader import DataLoader
import pdb
from dummy_data import append_label
from graph_loader import append_pos_label, train_loader, train_dataset, test_dataset
from sklearn.metrics import classification_report
from sklearn.metrics import matthews_corrcoef
import random
dataset = train_dataset
def goodness(data):
goodness = data.pow(2).mean(1)
return goodness
def loss_ff(x, positive):
threshold = 2
theta = threshold if positive else -threshold
out = -x if positive else x ### Loss is calculated different for positive and negative examples
loss = torch.log(1+torch.exp(out+theta)).mean()
# print("Loss: ",loss)
return loss
class GNN_FF(torch.nn.Module):
def __init__(self):
super().__init__()
self.layer1 = GCNConv(-1, 32)
self.layer2 = GCNConv(32, 2)
self.norm = torch.nn.LayerNorm(dataset.num_node_features, 32)
self.relu = torch.nn.ReLU()
self.goodness = goodness
self.loss_ff = loss_ff
self.layers = [self.layer1, self.layer2]
def forward(self, data):
x = data.x
edge_index = data.edge_index
g = []
for layer in self.layers:
x = x.detach()
x = layer(x, edge_index)
x = torch.nan_to_num(x)
x = self.relu(x)
g = goodness(x)
g += g
return x, torch.stack(tuple(g), 0).sum(0)
def train_ff(self, data, positive, optimizer):
x, edge_index = data.x, data.edge_index
for layer in self.layers:
x = x.detach()
x = layer(x, edge_index)
x = torch.nan_to_num(x)
x = self.relu(x)
out = goodness(x)
loss = loss_ff(out, positive)
loss = torch.tensor(loss, requires_grad=True) ### maybe False
print("Loss: ", loss)
optimizer.zero_grad()
loss.backward()
optimizer.step()
return loss
device='cpu'
# device = torch.device('cuda' if torch.cuda.is_available else 'cpu')
model = GNN_FF().to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
def train(model, train_loader, optimizer, device):
model.train()
for epoch in range(10):
for graph in train_loader:
for i in range(2):
if i == 0:
label = torch.Tensor([0, 1])
x = append_label(graph.x, graph.edge_index, label)
loss = model.train_ff(x, False, optimizer)
if i == 1:
label = torch.Tensor([1, 0])
x = append_label(graph.x, graph.edge_index, label)
loss = model.train_ff(x, True, optimizer)
train(model=model, train_loader=train_dataset, optimizer=optimizer, device=device)
def predict(model, pred_data):
sample = pred_data
g_for_label = []
gg = []
for i in range(2):
if i == 0:
label = torch.Tensor([0,1])
else:
label = torch.Tensor([1,0])
x = append_label(sample.x, sample.edge_index, label)
_, g = model.forward(x)
print("g:", g)
g_for_label += [g]
print("G_for_label:", g_for_label)
print("Total goodness: ", torch.stack(g_for_label, 0))#.argmax(0))
return torch.stack(g_for_label, 0).argmax(0)
def test(model, test_loader, device):
correct = 0
total = 0
pred_s = []
corr = []
model.eval()
with torch.no_grad():
for graph in test_loader:
x, label = graph.x, graph.label
pred = predict(model, graph)
print("Prediction: ", pred)
correct += (pred == label).sum().item()
print("Correct: ", correct)
total += label
print("Total: ", total)
acc = correct / total
pred_s.append(pred)
corr.append(label)
print("Pred_s: ",pred)
print("Corr: ", corr)
preds = []
for i in pred_s:
if i == 0:
score = [0, 1]
else:
score = [1, 0]
preds += score
print("Preds: ", preds)
corrects = []
for i in corr:
if i ==1:
s = [1, 0]
corrects += s
mcc = matthews_corrcoef(torch.tensor(corrects), preds)
print("Accuracy: ", acc)
print("MCC: ", mcc)
print("Scores: ", classification_report(corr, pred_s, labels=[1]))
return acc
test(model=model, test_loader=test_dataset, device=device)
