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

使用GATv2Conv构建图分类网络时出现维度错误的解决方法

解决GATv2Conv图分类中的维度不匹配错误

你在使用自定义GraphClassifier构建GATv2Conv网络进行图分类时,遇到维度错误,错误指向forward函数中的h = self.conv2(g, h)行,错误信息为:

RuntimeError: mat1 and mat2 shapes cannot be multiplied (7140x16 and 64x64)

模型代码

class GraphClassifier(nn.Module):
    def __init__(self, in_feats, hidden_size, num_classes):
        super(GraphClassifier, self).__init__()
        self.conv1 = GATv2Conv(in_feats, hidden_size, num_heads=4)
        self.conv2 = GATv2Conv(4*hidden_size, hidden_size, num_heads=4)
        self.conv3 = GATv2Conv(4*hidden_size, hidden_size, num_heads=4)
        self.conv4 = GATv2Conv(4*hidden_size, hidden_size, num_heads=1)
        self.classify = nn.Linear(hidden_size, num_classes)
        self.dropout = nn.Dropout(p=0.5)
        
    def forward(self, g, inputs):
        h = self.conv1(g, inputs)
        h = F.elu(h)
        h = self.dropout(h)
        
        h = self.conv2(g, h)
        h = F.elu(h)
        h = self.dropout(h)
        
        h = self.conv3(g, h)
        h = F.elu(h)
        h = self.dropout(h)
        
        h = self.conv4(g, h)
        h = F.elu(h)
        h = self.dropout(h)
        
        with g.local_scope():
            g.ndata['h'] = h
            # Calculate graph representation by max pooling readout.
            hg = dgl.max_nodes(g, 'h')
            return self.classify(hg)

训练代码

import torch.nn.functional as F
model = GraphClassifier(dataset.dim_nfeats, 16, dataset.gclasses)
        
opt = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=1e-4)

for epoch in range(400):
    
    model.train()               
    cumulative_loss_train = 0.0 
    
    for batched_graph, labels in trainloader:
        features = batched_graph.ndata['attr']
        
        logits = model(batched_graph, batched_graph.ndata["attr"])
                       
        loss = F.cross_entropy(logits, labels) 
        
        opt.zero_grad()         
        loss.backward()         
        opt.step()         
        
        cumulative_loss_train += loss.item()

错误原因

问题出在GATv2Conv的输出维度处理上:

  • 当设置num_heads=4时,GATv2Conv的输出形状为[num_nodes, num_heads, hidden_size](三维张量),而非预期的[num_nodes, num_heads*hidden_size](二维张量)。
  • conv2定义为GATv2Conv(4*hidden_size, hidden_size, num_heads=4),期望输入是每个节点64维(4*16),但conv1输出的三维张量直接传入时,实际被视为每个节点16维(仅取最后一维),导致维度不匹配,触发矩阵乘法错误。

解决方法

在每个GAT层的非线性激活后,将多头的特征维度合并为一维,即在每个F.elu(h)之后添加h = h.flatten(1),将三维张量[num_nodes, heads, hidden]转为二维张量[num_nodes, heads*hidden]。修改后的forward函数如下:

def forward(self, g, inputs):
    h = self.conv1(g, inputs)
    h = F.elu(h)
    h = h.flatten(1)  # 合并多头维度
    h = self.dropout(h)
    
    h = self.conv2(g, h)
    h = F.elu(h)
    h = h.flatten(1)  # 合并多头维度
    h = self.dropout(h)
    
    h = self.conv3(g, h)
    h = F.elu(h)
    h = h.flatten(1)  # 合并多头维度
    h = self.dropout(h)
    
    h = self.conv4(g, h)
    h = F.elu(h)
    h = h.flatten(1)  # 合并最后一层的多头维度(num_heads=1,不影响)
    h = self.dropout(h)
    
    with g.local_scope():
        g.ndata['h'] = h
        hg = dgl.max_nodes(g, 'h')
        return self.classify(hg)

最后一层conv4设置num_heads=1,合并后维度即为hidden_size,刚好匹配后续线性层classify的输入维度,无需额外调整。

内容的提问来源于stack exchange,提问作者Branco François

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 13:47:02