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

神经网络矩阵维度匹配错误求助:matmul运算维度不兼容问题

问题分析与解决方案

错误根源

你遇到的矩阵维度不匹配错误,核心原因是**backward函数没有获取到前向传播过程中生成的正确中间变量**:

  • 你在forward中计算了h1、t1、x这些关键变量,但调用backward时只传入了z和y,导致backward里引用的h1、t1、x不是当前样本对应的有效变量(可能是未定义的全局变量或旧值),维度自然不匹配。
  • 此外,代码还存在两个隐性问题:
    1. update函数修改全局权重时未声明global,导致实际修改的是局部变量,权重根本不会更新。
    2. 循环变量名重复(外层和内层都用i),会导致循环计数逻辑混乱。

修正步骤

1. 修正forward函数,返回必要的中间变量

需要把前向传播中的t1也返回,供反向传播计算ReLU导数使用:

def forward(x, y):
    t1 = x @ W1 + b1
    h1 = relu(t1)
    t2 = h1 @ W2 + b2
    z = softmax(t2)
    E = cross_entropy(z, y)
    return z, E, h1, t1  # 新增返回t1

2. 修改backward函数,传入所需参数

将h1、t1、x作为参数传入,确保使用当前样本的中间结果:

def backward(z, y, h1, t1, x):
    y_full = to_full(y, OUTPUT_DIM)
    dE_dt2 = z - y_full
    dE_dW2 = h1.T @ dE_dt2  # 此时h1是(1,25),h1.T是(25,1),dE_dt2是(1,10),矩阵乘法结果为(25,10),与W2维度匹配
    dE_db2 = dE_dt2
    dE_dh1 = dE_dt2 @ W2.T
    dE_dt1 = dE_dh1 * relu_deriv(t1)
    dE_dW1 = x.T @ dE_dt1
    dE_db1 = dE_dt1
    return dE_dW1, dE_db1, dE_dW2, dE_db2

3. 修正update函数,声明全局变量

添加global关键字,确保修改的是全局权重变量:

def update(dE_dW1, dE_db1, dE_dW2, dE_db2):
    global W1, b1, W2, b2
    W1 = W1 - ALPHA * dE_dW1
    b1 = b1 - ALPHA * dE_db1.mean(axis=0)  # 单样本训练时mean可不加,但批量训练时需要
    W2 = W2 - ALPHA * dE_dW2
    b2 = b2 - ALPHA * dE_db2.mean(axis=0)

4. 修正循环变量名,传递正确参数

把外层循环变量改为epoch,内层改为idx,并在调用backward时传入所有必要参数:

for epoch in range(NUM_EPOCHS):
    for idx in range(len(data)):
        x, y = data[idx]
        z, E, h1, t1 = forward(x, y)
        dE_dW1, dE_db1, dE_dW2, dE_db2 = backward(z, y, h1, t1, x)
        update(dE_dW1, dE_db1, dE_dW2, dE_db2)

额外检查点

确保你的辅助函数维度输出正确:

  • relu和relu_deriv应保持输入的维度(输入是(1,25),输出也应为(1,25))
  • softmax处理(1,10)的输入时,输出也应为(1,10)
  • to_full应将标签y转换为(1, OUTPUT_DIM)的one-hot向量

内容的提问来源于stack exchange,提问作者Александр Ищенко

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 14:38:15