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

PyTorch & DGL矩阵乘法报错:mat1与mat2无法相乘(1x4581和1x4581)

DGL GCN训练报错:RuntimeError: mat1 and mat2 shapes cannot be multiplied (1x4581 and 1x4581)

问题重现

使用DGL训练图卷积神经网络时触发矩阵维度不匹配错误,核心代码及报错如下:

训练代码

import dgl
from dgl.nn import GraphConv
import torch
import torch.nn as nn
import torch.nn.functional as F

class GCN(nn.Module):
    def __init__(self, in_feats, h_feats, num_classes):
        super(GCN, self).__init__()
        self.conv1 = GraphConv(in_feats, h_feats, allow_zero_in_degree=True)
        self.conv2 = GraphConv(h_feats, num_classes, allow_zero_in_degree=True)

    def forward(self, g, in_feat):
        h = self.conv1(g, in_feat)
        h = F.relu(h)
        h = self.conv2(g, h)
        g.ndata['h'] = h
        h_mean = dgl.mean_nodes(g, 'h')
        return h_mean

# 初始化模型与优化器
model = GCN(1, 4581, 1)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)

# 训练循环
for epoch in range(10):
    for batched_graph, labels in train_data:
        pred = model(batched_graph, batched_graph.ndata['cond'].float())
        loss = F.cross_entropy(pred, labels)
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

报错信息

RuntimeError: mat1 and mat2 shapes cannot be multiplied (1x4581 and 1x4581)

错误原因

  1. 任务与损失函数不匹配:F.cross_entropy要求分类任务的类别数≥2,但你设置num_classes=1,且用dgl.mean_nodes输出[batch_size,1]的结果,损失函数内部处理时会引发维度异常。
  2. 中间层维度设置不合理:h_feats=4581远大于输入特征维度1,不仅会导致计算量爆炸,还容易在单节点图的batch中触发矩阵相乘维度错误——当batch中存在仅1个节点的图时,第一层输出为[1,4581],第二层GraphConv的权重矩阵为[4581,1],若因损失函数的错误处理导致维度错位,就会出现1x4581与1x4581相乘的非法操作。
  3. 输入特征维度可能缺失:若batched_graph.ndata['cond']的形状是[total_nodes]而非[total_nodes,1],会被DGL误判为节点数=total_nodes、特征维度=1,但部分场景下会引发维度传递错误。

解决方案

1. 匹配任务与损失函数

  • 回归任务:替换损失函数为均方误差损失,同时保持num_classes=1:
    # 训练循环内修改损失计算
    pred = model(batched_graph, batched_graph.ndata['cond'].float())
    loss = F.mse_loss(pred.squeeze(), labels.float())
    
  • 二分类任务:修改模型类别数为2,若为图级分类可保留mean_nodes,节点分类则直接返回节点特征:
    # 修改模型初始化
    model = GCN(1, 64, 2)
    
    # 若为节点分类,修改forward函数
    def forward(self, g, in_feat):
        h = self.conv1(g, in_feat)
        h = F.relu(h)
        h = self.conv2(g, h)
        return h
    
    # 损失计算
    pred = model(batched_graph, batched_graph.ndata['cond'].float())
    loss = F.cross_entropy(pred, labels)
    

2. 调整中间层维度

将h_feats改为合理的小数值(如64、128),避免计算资源浪费与维度异常:

model = GCN(1, 64, 1)  # 回归任务
# 或
model = GCN(1, 128, 2) # 二分类任务

3. 确保输入特征维度正确

检查并修正输入特征的形状,确保为[total_nodes, 1]:

in_feat = batched_graph.ndata['cond'].float().unsqueeze(1)
pred = model(batched_graph, in_feat)

内容的提问来源于stack exchange,提问作者77 ChickenNug

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 16:43:12