You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.01 11:09:04