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)
错误原因
- 任务与损失函数不匹配:
F.cross_entropy要求分类任务的类别数≥2,但你设置num_classes=1,且用dgl.mean_nodes输出[batch_size,1]的结果,损失函数内部处理时会引发维度异常。 - 中间层维度设置不合理:
h_feats=4581远大于输入特征维度1,不仅会导致计算量爆炸,还容易在单节点图的batch中触发矩阵相乘维度错误——当batch中存在仅1个节点的图时,第一层输出为[1,4581],第二层GraphConv的权重矩阵为[4581,1],若因损失函数的错误处理导致维度错位,就会出现1x4581与1x4581相乘的非法操作。 - 输入特征维度可能缺失:若
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
相关产品推荐
相关产品推荐

