PyTorch Geometric实现GNN运行报TypeError参数错误如何解决
错误原因
- 你定义的
Net类的forward方法缺失Python类实例方法必备的第一个位置参数self。所有Python类的实例方法在定义时都需要将指代实例本身的self作为首个参数,PyTorch的nn.Module子类的forward方法也遵循该语法规则。 - PyTorch执行
model()调用时,会自动把模型实例作为第一个参数传入forward方法,你定义的forward没有预留参数位,因此触发参数数量不匹配的TypeError。
修复方案
基础修复(兼容你现有调用逻辑)
直接给forward方法补充self参数即可正常运行:
class Net(torch.nn.Module): def __init__(self): super().__init__() self.conv = SAGEConv(dataset.num_features, dataset.num_classes, aggr="max") # max, mean, add ...) # 补充self参数 def forward(self): x = self.conv(data.x, data.edge_index) return F.log_softmax(x, dim=1)
优化写法(推荐)
不建议在forward内部直接依赖全局变量data,改为将输入作为参数传入,代码通用性更强:
模型定义修改
class Net(torch.nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv = SAGEConv(in_channels, out_channels, aggr="max") def forward(self, x, edge_index): x = self.conv(x, edge_index) return F.log_softmax(x, dim=1)
调用逻辑同步修改
def train(): model.train() optimizer.zero_grad() # 传入模型所需的输入参数 F.nll_loss(model(data.x, data.edge_index)[data.train_mask], data.y[data.train_mask]).backward() optimizer.step() def test(): model.eval() logits = model(data.x, data.edge_index) accs = [] for _, mask in data('train_mask', 'val_mask', 'test_mask'): pred = logits[mask].max(1)[1] acc = pred.eq(data.y[mask]).sum().item() / mask.sum().item() accs.append(acc) return accs
内容的提问来源于stack exchange,提问作者James Arten
相关产品推荐
相关产品推荐

