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

自建三层神经网络MNIST预测异常,求问题排查建议

三层神经网络MNIST预测异常排查求助

自行搭建三层神经网络用于MNIST数据集预测,参考在线代码并自主实现部分逻辑。代码无语法错误,但训练后网络无论输入何种样本,总有1-2个类别概率极高,预测结果全错。以下是完整代码及测试输出:

import numpy as np
from PIL import Image
import os
np.set_printoptions(formatter={'float_kind':'{:f}'.format})
def init_setup():
    #three layers perception
    w1=np.random.randn(10,784)-0.8
    b1=np.random.rand(10,1)-0.8
    #second layer
    w2=np.random.randn(10,10)-0.8
    b2=np.random.randn(10,1)-0.8
    #third layer
    w3=np.random.randn(10,10)-0.8
    b3=np.random.randn(10,1)-0.8
    return w1,b1,w2,b2,w3,b3
def activate(A):
    # use ReLU function as the activation function
    Z=np.maximum(0,A)
    return Z
def softmax(Z):
    return np.exp(Z)/np.sum(np.exp(Z))

def forward_propagation(A,w1,b1,w2,b2,w3,b3):
    # input A :(784,1)-> A1: (10,1) ->A2: (10,1) -> prob: (10,1)
    z1=w1@A+b1
    A1=activate(z1)
    z2=w2@A1+b2
    A2=activate(z2)
    z3=w3@A2+b3
    prob=softmax(z3)

    return z1,A1,z2,A2,z3,prob
def one_hot(Y:np.ndarray)->np.ndarray:

    one_hot=np.zeros((10, 1)).astype(int)
    
    one_hot[Y]=1
    return one_hot

def back_propagation(A,z1,A1:np.ndarray,z2,A2:np.ndarray,z3,prob,w1,w2:np.ndarray,w3,Y:np.ndarray,lr:float):

    m=1/Y.size

    dz3=prob-Y 

    dw3=m*dz3@A2.T

    db3= dz3
    dz2=ReLU_deriv(z2)*w3.T@dz3
    dw2 =  dz2@A1.T
    db2 =  dz2
    dz1=ReLU_deriv(z1)*w2.T@dz2
    dw1 = dz1@A.T
    db1 =  dz1
    return db1,dw1,dw2,db2,dw3,db3
def ReLU_deriv(Z):
    Z[Z>0]=1
    Z[Z<=0]=0
    return Z 
def step(lr,w1,b1,w2,b2,w3,b3,dw1,db1,dw2,db2,dw3,db3):
    w1 = w1 - lr * dw1

    b1 = b1 - lr * db1    
    w2 = w2 - lr * dw2  
    b2 = b2 - lr * db2
    w3 = w3 - lr * dw3 
    b3 = b3 - lr * db3       
    return w1,b1,w2,b2,w3,b3

整合训练函数

def learn():
    lr=0.5
    dir=r'C:\Users\Desktop\MNIST - JPG - training\{}'
    w1,b1,w2,b2,w3,b3=init_setup()
    for e in range(10):
        if e%3 == 0:
            lr=lr/10
        for num in range(10):
            Y=one_hot(num)
            # print(Y)
            path=dir.format(str(num))
            for i in os.listdir(path):
                img=Image.open(path+'\\'+i)
                A=np.asarray(img)
                A=A.reshape(-1,1) 
                z1,A1,z2,A2,z3,prob=forward_propagation(A,w1,b1,w2,b2,w3,b3)
                # print('loss='+str(np.sum(np.abs(Y-prob))))
                db1,dw1,dw2,db2,dw3,db3=back_propagation(A,z1,A1,z2,A2,z3,prob,w1,w2,w3,Y,lr)
                w1,b1,w2,b2,w3,b3=step(lr,w1,b1,w2,b2,w3,b3,dw1,db1,dw2,db2,dw3,db3)
    return  w1,b1,w2,b2,w3,b3
optimize_params=learn()
w1,b1,w2,b2,w3,b3=optimize_params
img=Image.open(r'C:\Users\Desktop\MNIST - JPG - training\2\5.jpg')
A=np.asarray(img)
A=A.reshape(-1,1)
z1,A1,z2,A2,z3,prob=forward_propagation(A,w1,b1,w2,b2,w3,b3)
print(prob)
print(np.argmax(prob))

测试输出

>>>[[0.040939]
    [0.048695]
    [0.048555]
    [0.054962]
    [0.060614]
    [0.066957]
    [0.086470]
    [0.117370]
    [0.163163]
    [0.312274]]
>>>9

真实标签为2,但类别2的概率极低,预测完全错误,恳请提供排查方向。


排查方向
  • 数据未归一化:MNIST图像像素值范围是0-255,直接输入会导致权重更新时数值过大,引发梯度爆炸或网络偏向。需将像素值除以255缩放到0-1区间。
  • 权重初始化错误:当前init_setup中用np.random.randn(...) - 0.8,会让初始权重整体偏向负数,导致ReLU激活后大量神经元输出为0,网络无法有效学习。建议使用He初始化(针对ReLU):w1 = np.random.randn(10,784) * np.sqrt(2/784),偏置初始化为0或小值,不要整体偏移-0.8。
  • 反向传播梯度未正确平均:当前是单样本SGD,但back_propagation中dw2 = dz2@A1.T等操作未除以样本数m,会导致梯度更新幅度过大。需统一对所有权重梯度除以m,保证更新幅度合理。
  • ReLU导数的原地修改问题:ReLU_deriv函数中直接修改输入ZZ[Z>0]=1,会污染前向传播的中间结果,应该创建副本:return np.where(Z>0, 1, 0)。
  • 学习率调整过于激进:训练初始lr=0.5,每3个epoch直接除以10,e=0时lr就变成0.05,后续快速降到极小值,网络还未充分学习就无法有效更新权重。建议调整策略,比如初始lr=0.01,或每5个epoch将学习率乘以0.5。
  • Softmax数值稳定性问题:当前softmax函数直接计算np.exp(Z)/np.sum(np.exp(Z)),当Z中数值较大时,exp会溢出导致结果异常。需改进为:exp_Z = np.exp(Z - np.max(Z)),再除以sum(exp_Z),避免数值溢出。
  • 损失函数监控:取消print('loss='+str(np.sum(np.abs(Y-prob))))的注释,观察训练过程中损失是否下降。如果损失不下降或波动极大,说明梯度更新存在问题;如果损失一直很高,说明网络无法有效学习。
  • 网络结构容量不足:三层网络每层仅10个神经元,表达能力有限。可尝试增加中间层神经元数量,比如第一层用64或128个神经元,提升网络拟合能力。

内容的提问来源于stack exchange,提问作者yao3414

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 00:21:41