神经网络训练报错:矩阵维度不匹配问题求助
问题描述
搭建一个包含5个输入、4个隐藏层节点、1个输出的神经网络,设置学习率0.2,误差阈值0.2,数据从CSV文件(实际使用泰坦尼克数据集)读取。训练过程中出现如下错误:
ValueError: shapes (1,6) and (5,5) not aligned: 6 (dim 1) != 5 (dim 0)
定位到代码中hidden_in = np.dot(inputs, w1)行,怀疑是权重矩阵与输入的矩阵乘法维度不匹配导致问题,请求解决。相关代码如下:
def train(inputs_list, w1, w2, w3, targets_list, lr, error): era = 0 list_error = [] global_error = float('inf') while global_error > error: local_error = np.array([]) for i, inputs in enumerate(inputs_list): inputs = np.array(inputs, ndmin=2) targets = np.array(targets_list[i], ndmin=2) # 前向传播 hidden_in = np.dot(inputs, w1) hidden_out = f(hidden_in) hidden_out = np.array(np.insert(hidden_out, 0, [1]), ndmin=2) hidden_in2 = np.dot(hidden_out, w2) hidden_out2 = f(hidden_in2) hidden_out2 = np.array(np.insert(hidden_out2, 0, [1]), ndmin=2) final_in = np.dot(hidden_out2, w3) final_out = final_in # 误差计算 output_error = targets - final_out hidden_error2 = np.dot(output_error, w3.T) hidden_error = np.dot(hidden_error2[:, 1:], w2.T) local_error = np.append(local_error, output_error) # 反向传播更新权重 w3 += lr * output_error * hidden_out2.T w2 += lr * hidden_error2[:, 1:] * f1(hidden_out2[:, 1:]) * hidden_out.T w1 += lr * hidden_error[:, 1:] * f1(hidden_out[:, 1:]) * inputs.T global_error = abs(np.mean(local_error)) era += 1 list_error.append(global_error) if era > 1000: print('gl=', global_error) break print(global_error) return w1, w2, w3, era, list_error def query(inputs_list, w1, w2, w3): final_out = np.array([]) for i, inputs in enumerate(inputs_list): inputs = np.array(inputs, ndmin=2) hidden_in = np.dot(inputs, w1) hidden_out = f(hidden_in) hidden_out = np.array(np.insert(hidden_out, 0, [1]), ndmin=2) hidden_in2 = np.dot(hidden_out, w2) hidden_out2 = f(hidden_in2) hidden_out2 = np.array(np.insert(hidden_out2, 0, [1]), ndmin=2) final_in = np.dot(hidden_out2, w3) final_out = np.append(final_out, final_in) return np.around(final_out) # 数据读取与预处理 import pandas as pd import numpy as np data_titanic = "titanic_dataset.csv" data = pd.read_csv(data_titanic) target_data = data['Survived'].values data = data.drop('Survived', 1).values # 训练集与测试集划分 inputs = data[0:600] inputs = np.c_[np.ones(600), inputs] # 添加偏置项 targets = target_data[0:600] test = data[600:714] test = np.c_[np.ones(114), test] targets_test = target_data[600:714] # 参数设置 lr = 0.2 eps = 0.2 input_layer = 5 # 此处定义错误 hidden_layer = 4 hidden_layer2 = 2 output_layer = 1 # 权重初始化(假设init_weight函数按输入层和隐藏层维度生成权重矩阵) def init_weight(input_dim, hidden1_dim, hidden2_dim, output_dim): w1 = np.random.randn(input_dim, hidden1_dim) w2 = np.random.randn(hidden1_dim + 1, hidden2_dim) # +1对应隐藏层的偏置项 w3 = np.random.randn(hidden2_dim + 1, output_dim) return w1, w2, w3 w1, w2, w3 = init_weight(input_layer, hidden_layer, hidden_layer2, output_layer) w1, w2, w3, era, lst = train(inputs, w1, w2, w3, targets, lr, eps)
问题分析与解决
核心原因
错误提示明确说明输入矩阵维度(1,6)与权重矩阵w1的(5,4)不兼容,矩阵乘法要求左矩阵的列数等于右矩阵的行数。问题出在两处:
- 输入层维度定义错误:预处理时给每个输入样本添加了1列偏置项(
np.c_[np.ones(600), inputs]),导致输入样本的维度从5变为6,但代码中input_layer仍设为5。 - 权重矩阵w1的初始化维度不匹配:
init_weight函数根据input_layer=5生成形状为(5,4)的w1,而实际输入是6维,导致矩阵乘法时维度冲突。
修正步骤
- 修正输入层维度参数:将
input_layer改为6,对应添加偏置后的输入维度:input_layer = 6 # 原输入5维 + 1维偏置 - 确保权重矩阵维度兼容:
- w1的形状应为
(输入维度, 第一个隐藏层节点数),即(6,4),对应输入6维到4个隐藏节点的映射。 - 修改
input_layer后,init_weight函数会自动生成正确维度的w1。
- w1的形状应为
- 验证其他权重矩阵的维度:
- w2的形状应为
(第一个隐藏层节点数+1, 第二个隐藏层节点数),即(4+1,2)=(5,2),对应第一个隐藏层输出(含偏置)到第二个隐藏层的映射。 - w3的形状应为
(第二个隐藏层节点数+1, 输出层节点数),即(2+1,1)=(3,1),对应第二个隐藏层输出(含偏置)到输出层的映射。
- w2的形状应为
修正后的关键代码片段
# 参数设置修正 input_layer = 6 # 修正为添加偏置后的输入维度 hidden_layer = 4 hidden_layer2 = 2 output_layer = 1 # 权重初始化函数(确保维度正确) def init_weight(input_dim, hidden1_dim, hidden2_dim, output_dim): w1 = np.random.randn(input_dim, hidden1_dim) # (6,4) w2 = np.random.randn(hidden1_dim + 1, hidden2_dim) # (5,2) w3 = np.random.randn(hidden2_dim + 1, output_dim) # (3,1) return w1, w2, w3 w1, w2, w3 = init_weight(input_layer, hidden_layer, hidden_layer2, output_layer) w1, w2, w3, era, lst = train(inputs, w1, w2, w3, targets, lr, eps)
额外注意事项
- 若后续调整网络结构(比如增加更多隐藏层),需始终保证每一层的输入维度与对应权重矩阵的行数一致。
- 偏置项的处理要统一:要么在输入层提前添加,要么在每个隐藏层的前向传播中单独处理,避免重复添加导致维度混乱。
内容的提问来源于stack exchange,提问作者xor
相关产品推荐
相关产品推荐

