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

神经网络训练matmul维度不匹配错误:代码报错求助

解决神经网络训练中matmul矩阵乘法维度不匹配问题

错误信息

matmul: Input operand 1 has a mismatch in its core dimension 0, 
with gufunc signature (n?,k),(k,m?)->(n?,m?) (size 4 is different from 2)

问题根源

  1. 拓扑结构与输入特征维度不匹配:拓扑结构topology = [p, 4, 8, 16, 8, 4, 1]中的第一个元素p代表输入特征维度,但实际输入数据X的特征数为2,而p被设置为4,导致第一层权重矩阵的输入维度(4)与X的特征维度(2)不匹配,矩阵乘法无法执行。
  2. 代码索引错误:print(out[-1[1]])的索引写法错误,正确应为print(out[-1][1]),否则会触发额外的索引错误。

解决步骤

  • 对齐输入维度与拓扑结构:确认输入数据X的形状(X.shape),X应为(样本数量, 输入特征数),将拓扑结构的第一个元素p设置为与输入特征数一致的值。
  • 修正索引错误:将out[-1[1]]改为out[-1][1]。
  • 验证权重矩阵维度:确保create_nn函数生成的每层权重矩阵W的形状为(前一层神经元数, 当前层神经元数),比如第一层W的形状应为(p, 4),第二层为(4, 8),以此类推。

修正后的完整代码示例

import numpy as np

# 定义sigmoid激活函数(含导数)
def sigm(z):
    return 1/(1+np.exp(-z)), z*(1-z)

# 定义神经网络层类
class Layer:
    def __init__(self, input_dim, output_dim, act_f):
        self.W = np.random.randn(input_dim, output_dim) * 0.1  # 初始化权重
        self.b = np.zeros((1, output_dim))                     # 初始化偏置
        self.act_f = act_f

# 根据拓扑结构创建神经网络
def create_nn(topology, act_f):
    neural_net = []
    for idx in range(len(topology)-1):
        neural_net.append(Layer(topology[idx], topology[idx+1], act_f))
    return neural_net

# 输入特征数设为2(匹配错误提示中的size 2)
p = 2
topology = [p, 4, 8, 16, 8, 4, 1]

neural_net = create_nn(topology, sigm)

# L2损失函数(含导数)
l2_cost = (lambda Yp, Yr: np.mean((Yp - Yr) ** 2),
           lambda Yp, Yr: (Yp - Yr))

def train(neural_net, X, Y, l2_cost, lr=0.5):
    out = [(None, X)]
    
    # 前向传播
    for layer in neural_net:
        z = out[-1][1] @ layer.W + layer.b
        a = layer.act_f[0](z)
        out.append((z, a))
    
    # 输出最后一层的激活值(修正索引错误)
    print(out[-1][1])

# 生成示例输入数据:5个样本,每个样本2个特征
X = np.random.randn(5, 2)
Y = np.random.randn(5, 1)

# 启动训练
train(neural_net, X, Y, l2_cost, 0.5)

内容的提问来源于stack exchange,提问作者JONATHAN . MALDONADO MIRANDA

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 20:48:45