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

基于Numpy从零搭建的MNIST分类神经网络无法学习的问题求助

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 19:08:11