基于Numpy从零搭建的MNIST分类神经网络无法学习的问题求助
我用Python和Numpy从零搭建了MNIST分类神经网络,代码能运行但无法有效学习:初始准确率约12%,20轮后升至14%,40轮回落至12%,150轮仍无改善,推测反向传播存在问题。我参考的教程以行表示特征、列表示样本,但我用的是列表示特征、行表示样本(输入数据shape为(60000,784)),反向传播时转置数组适配算法,怀疑维度处理错误是问题根源。
加载数据
(x_train, y_train), (x_test, y_test) = mnist.load_data() x_train, x_test = x_train / 255, x_test / 255 x_train, x_test = x_train.reshape(len(x_train), 28 * 28), x_test.reshape(len(x_test), 28 * 28) print(x_train.shape) # (60000, 784) print(x_test.shape) # (10000, 784)
模型核心代码
W1 = np.random.randn(784, 10) b1 = np.random.randn(10) W2 = np.random.randn(10, 10) b2 = np.random.randn(10) def relu(x, dir=False): if dir: return x > 0 return np.maximum(x, 0) def softmax(x): e_x = np.exp(x - np.max(x)) return e_x / e_x.sum(axis=1, keepdims = True) def one_hot_encode(y): y_hot = np.zeros(shape=(len(y), 10)) for i in range(len(y)): y_hot[i][y[i]] = 1 return y_hot def loss_function(predictions, true): return predictions - true def predict(x): Z1 = x.dot(W1) + b1 A1 = relu(Z1) Z2 = A1.dot(W2) + b2 A2 = softmax(Z2) # The final prediction is A2 at index 3 or -1: return Z1, A1, Z2, A2 def get_accuracy(predictions, Y): guesses = predictions.argmax(axis=1) average = 0 i = 0 while i < len(guesses): if guesses[i] == Y[i]: average += 1 i += 1 percent = (average / len(guesses)) * 100 return percent def train(data, labels, epochs=40, learning_rate=0.1): for i in range(epochs): labels_one_hot = one_hot_encode(labels) # Forward: m = len(labels_one_hot) Z1, A1, Z2, A2 = predict(data) # I think the error is in this chunk: # backwards pass: dZ2 = A2 - labels_one_hot dW2 = 1 / m * dZ2.T.dot(A1) db2 = 1 / m * np.sum(dZ2, axis=1) dZ1 = W2.dot(dZ2.T).T * relu(Z1, dir=True) dW1 = 1 / m * dZ1.T.dot(data) db1 = 1 / m * np.sum(dZ1) # Update parameters: update(learning_rate, dW1, db1, dW2, db2) print("Iteration: ", i + 1) predictions = predict(data)[-1] # item at -1 is the final prediction. print(get_accuracy(predictions, labels)) def update(learning_rate, dW1, db1, dW2, db2): global W1, b1, W2, b2 W1 = W1 - learning_rate * dW1.T # I have to transpose it here. b1 = b1 - learning_rate * db1 W2 = W2 - learning_rate * dW2 b2 = b2 - learning_rate * db2 train(x_train, y_train) predictions = predict(x_test)[-1] print(get_accuracy(predictions, y_test)) # The result is about 11.5% accuracy.
问题排查与修复
你的反向传播和参数更新中的维度处理存在多处错误,以下是具体问题和修正方案:
1. dW2 计算错误
原代码中 dW2 = 1/m * dZ2.T.dot(A1),但根据你的数据维度(样本在行,特征在列),正确的计算应该是A1的转置点乘dZ2,保证dW2的shape与W2(10,10)一致。
修正:
dW2 = 1 / m * A1.T.dot(dZ2)
2. db2 计算错误
原代码中db2 = 1/m * np.sum(dZ2, axis=1),axis=1会按行求和得到(60000,)的数组,但b2是(10,)的向量,应该按列求和(axis=0)。
修正:
db2 = 1 / m * np.sum(dZ2, axis=0)
3. dZ1 计算错误
原代码中dZ1 = W2.dot(dZ2.T).T * relu(Z1, dir=True),正确的矩阵乘法应该是dZ2点乘W2的转置,保证dZ1与Z1(60000,10)维度一致。
修正:
dZ1 = dZ2.dot(W2.T) * relu(Z1, dir=True)
4. dW1 计算错误
原代码中dW1 = 1/m * dZ1.T.dot(data),正确的计算应该是data的转置点乘dZ1,保证dW1的shape与W1(784,10)一致。
修正:
dW1 = 1 / m * data.T.dot(dZ1)
5. db1 计算错误
原代码中db1 = 1/m * np.sum(dZ1),会把所有元素求和得到一个标量,但b1是(10,)的向量,应该按列求和(axis=0)。
修正:
db1 = 1 / m * np.sum(dZ1, axis=0)
6. update 函数中W1的转置冗余
原代码中W1 = W1 - learning_rate * dW1.T,修正后的dW1已经是(784,10)的正确shape,不需要再转置。
修正:
W1 = W1 - learning_rate * dW1
修正后的完整反向传播与更新代码
def train(data, labels, epochs=40, learning_rate=0.1): for i in range(epochs): labels_one_hot = one_hot_encode(labels) # Forward: m = len(labels_one_hot) Z1, A1, Z2, A2 = predict(data) # 修正后的反向传播: dZ2 = A2 - labels_one_hot dW2 = 1 / m * A1.T.dot(dZ2) db2 = 1 / m * np.sum(dZ2, axis=0) dZ1 = dZ2.dot(W2.T) * relu(Z1, dir=True) dW1 = 1 / m * data.T.dot(dZ1) db1 = 1 / m * np.sum(dZ1, axis=0) # Update parameters: update(learning_rate, dW1, db1, dW2, db2) print("Iteration: ", i + 1) predictions = predict(data)[-1] print(get_accuracy(predictions, labels)) def update(learning_rate, dW1, db1, dW2, db2): global W1, b1, W2, b2 W1 = W1 - learning_rate * dW1 b1 = b1 - learning_rate * db1 W2 = W2 - learning_rate * dW2 b2 = b2 - learning_rate * db2
额外优化建议
- 初始化权重时,建议使用He初始化(针对ReLU激活),避免梯度消失/爆炸:
W1 = np.random.randn(784,10) * np.sqrt(2/784),W2 = np.random.randn(10,10) * np.sqrt(2/10) - 学习率可以适当调低(比如0.01),避免参数更新幅度过大导致震荡
- 加入小批量梯度下降(Mini-batch GD),比全量梯度下降更稳定高效
内容的提问来源于stack exchange,提问作者buzzbuzz20xx

