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

训练简单MLP网络遇矩阵相乘报错:维度匹配疑惑求助

矩阵乘法维度不匹配报错排查

我正在处理的网络张量形式如下:

tensor([[0.],
        [0.],
        [1.],
        ...,
        [0.],
        [1.],
        [1.]])

运行代码时触发以下错误:

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

按数学规则,(m×n)与(p×m)的维度应该可以匹配相乘,但实际却报错了,想排查问题出在哪里。

我的训练代码如下:

def train_epoch_sparse(model, optimizer, device, graph, train_edges, batch_size, epoch, monet_pseudo=None):

    model.train()
    
    train_edges = train_edges.to(device)
    
    total_loss = total_examples = 0
    for perm in tqdm(DataLoader(range(train_edges.size(0)), batch_size, shuffle=True)):

        optimizer.zero_grad()

        graph = graph.to(device)
        x = graph.ndata['h'].to(device).float()
        e = graph.edata['h'].to(device).float()

        if monet_pseudo is not None:
            # Assign e as pre-computed pesudo edges for MoNet
            e = monet_pseudo.to(device)
        h = model(graph, x, e)
        # Positive samples
        edge = train_edges[perm].t()
        pos_out = model.edge_predictor( h[edge[0]], h[edge[1]] )
        # Just do some trivial random sampling
        edge = torch.randint(0, x.size(0), edge.size(), dtype=torch.long, device=x.device)

        neg_out = model.edge_predictor( h[edge[0]], h[edge[1]] )
        
        loss = model.loss(pos_out, neg_out)

        loss.backward()
        optimizer.step()

        num_examples = pos_out.size(0)
        total_loss += loss.detach().item() * num_examples
        total_examples += num_examples

    return total_loss/total_examples, optimizer

可能的错误原因

  • 矩阵乘法维度顺序错误:报错中的两个矩阵是4267x4267和1x4267,矩阵乘法要求第一个矩阵的列数等于第二个矩阵的行数,这里显然不满足。你可能误判了第二个矩阵的形状,它实际是1x4267而非预期的4267x1,导致维度不兼容。
  • edge_predictor实现逻辑问题:检查edge_predictor的代码,它接收两个节点特征h[edge[0]]和h[edge[1]],如果是做拼接后过线性层或直接内积,要确认特征维度是否正确。比如若h的形状是[N, D],则两个输入应该是[B, D],若其中一个被错误转置为[D, B],会直接引发后续矩阵乘法维度错误。
  • 节点特征h的维度异常:查看model(graph, x, e)返回的h的形状,正常应为[节点数, 特征维度],如果特征维度被错误设置为节点数(比如4267),就会出现4267x4267的特征矩阵,再和1x4267的矩阵相乘必然报错。
  • 张量转置/索引操作错误:代码中edge = train_edges[perm].t()转置后的形状是否符合预期?若train_edges[perm]是[B, 2],转置后是[2, B],此时h[edge[0]]应为[B, D],如果h维度异常或转置操作有误,会连锁引发后续维度问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 03:31:06