MLP前向传播维度不匹配报错:如何解决形状对齐问题?
问题解决方案
错误根源
报错的核心是前向传播中的嵌套循环逻辑错误:原代码外层遍历权重,内层遍历所有偏置,导致每处理一个权重时,会和所有偏置依次计算。第一次计算后激活值维度变为(5,),第二次内层循环时用这个(5,)去乘第一个权重(2,5),就出现了维度不匹配(5≠2)的问题。此外,前向传播的返回值也错误地返回了初始输入,而非最后一层输出。
修正后的代码
修正MLP类的关键方法
import numpy as np from random import random class MLP(object): def __init__(self, num_inputs=3, hidden_layers=[3, 3], num_outputs=2): self.num_inputs = num_inputs self.hidden_layers = hidden_layers self.num_outputs = num_outputs layers = [num_inputs] + hidden_layers + [num_outputs] weights = [] bias = [] for i in range(len(layers) - 1): w = np.random.rand(layers[i], layers[i + 1]) b = np.random.randn(layers[i+1]).reshape(1, layers[i+1]) weights.append(w) bias.append(b) self.weights = weights self.bias = bias activations = [] for i in range(len(layers)): a = np.zeros(layers[i]) activations.append(a) self.activations = activations def forward_propagate(self, inputs): activations = inputs # 确保输入为二维数组,避免广播问题(可选,但更稳妥) if activations.ndim == 1: activations = activations.reshape(1, -1) self.activations[0] = activations # 并行遍历对应的权重和偏置,而非嵌套循环 for i, (w, b) in enumerate(zip(self.weights, self.bias)): net_inputs = np.dot(activations, w) + b activations = self._sigmoid(net_inputs) self.activations[i + 1] = activations # 返回最后一层的输出,转为一维匹配目标维度 return activations.flatten() def train(self, inputs, targets, epochs, learning_rate): for i in range(epochs): sum_errors = 0 for j, input in enumerate(inputs): target = targets[j] output = self.forward_propagate(input) # 计算均方误差,后续可扩展反向传播逻辑 sum_errors += np.mean((output - target)**2) print(f"Epoch {i+1}, Loss: {sum_errors / len(inputs)}") def _sigmoid(self, x): y = 1.0 / (1 + np.exp(-x)) return y
测试代码(保持原有逻辑)
items = np.array([[random()/2 for _ in range(2)] for _ in range(1000)]) targets = np.array([[i[0] + i[1]] for i in items]) mlp = MLP(2, [5], 1) mlp.train(items, targets, 2, 0.1)
关键修正点
- 移除嵌套循环:使用
zip(self.weights, self.bias)同时遍历对应的权重和偏置,保证每一层的权重只和对应层的偏置配合计算。 - 修正返回值:前向传播返回最后一层的激活值,通过
flatten()转为一维,匹配目标数据的维度。 - 输入维度处理:可选地将一维输入转为二维数组,避免后续矩阵运算中的广播歧义。
- 添加损失打印:在train方法中添加了损失计算和打印,方便观察训练过程。
内容的提问来源于stack exchange,提问作者Al.Vioky
相关产品推荐
相关产品推荐

