ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

大权重初始化引发梯度爆炸:NumPy手写MNIST从NaN到97%准确率

大权重初始化引发梯度爆炸:NumPy手写MNIST从NaN到97%准确率 最近在用纯 NumPy 手写 MNIST 分类任务时我特意把权重初始化放大了一圈结果训练不到十个 epoch损失值直接跳到 NaN整个网络瞬间罢工。后来把初始化改小同样的结构、同样的数据准确率轻松跑到 97% 以上。这个对比让我意识到很多人学深度学习时容易忽略“初始化”这个看起来不起眼的细节但它在真实的数值计算中常常是决定训练能否进行下去的钥匙。这篇是“NumPy 上理解深度学习计算过程”系列的第二篇重点聊大权重初始化问题。我会从数学原理出发用纯 NumPy 搭建一个两层神经网络在 MNIST 数据集上复现大权重导致的梯度爆炸、激活饱和、训练崩溃等现象再给出可落地的初始化方案和排查技巧。适合正在手写反向传播、想从底层理解深度学习机制的同学也适合那些用框架训练时偶尔遇到 NaN 却不知道原因的人。1. 为什么要盯上“大权重初始化”1.1 看似无害的一行代码却是崩溃的源头在 NumPy 里初始化权重最顺手的方式就是np.random.randn(784, 100)。这句代码本身没有错但randn生成的是标准正态分布标准差为 1。对于 MNIST 这种输入维度为 784 的数据前向传播时第一个隐藏层的加权输入z1 x.dot(W1) b1中z1的每个分量是 784 个随机变量的和。假设输入x已经归一化到 [0,1]那么z1的标准差大约是sqrt(784) 28也就是说z1的数值范围很容易到 [-100, 100] 甚至更大。你可以想象一下784 个独立的高斯随机数相加它的标准差会随着维度增加而变大。这带来的直接后果是经过 sigmoid 激活函数后绝大多数神经元的输出会死死贴在 0 或 1 附近。sigmoid 函数在饱和区的导数接近 0这意味着反向传播时梯度会被“吃掉”靠近输入层的权重几乎得不到有效更新。而靠近输出层的那部分由于 logits 巨大softmax 的输出会变得极度尖锐交叉熵损失可能直接溢出NaN 就这么出现了。1.2 初始化问题在 MNIST 上的典型表现我在复现时用的是一个两层全连接网络输入层 784隐藏层 100输出层 10。如果按np.random.randn来初始化训练过程中的现象非常典型第一轮 loss 不是 0.69 这种“均匀分布”的合理值而是几十甚至上百第二轮开始出现nan。此时如果你打印z1的值会发现里面有大量正负 100 以上的数字激活值a1有近一半是 0另一半是 1整个网络的梯度信息几乎断层。很多初学者会误以为这是学习率太大造成的于是疯狂调小学习率。但你把学习率从 0.1 降到 0.0001发现 loss 虽然不崩了却像冻住了一样几乎不下降。原因很简单sigmoid 饱和导致的梯度消失不是学习率能救回来的。学习率只控制更新步长而梯度本身已经趋近于 0再怎么调小步长也没有意义。所以我想说遇到训练 NaN先别急着调学习率先检查初始化尺度是不是已经远超合理范围。2. 大权重初始化问题背后的计算本质2.1 前向传播中的数值放大要理解大权重为什么危险核心是看数据在网络中传播时方差的变化。假设输入x的每个分量均值为 0方差为 1做了标准化权重w也服从均值为 0、方差为Var[w]的分布那么某一层神经元的加权和z w1*x1 w2*x2 ... wn*xn的方差为Var[z] n * Var[w] * Var[x]如果w的标准差为 1n 784那么Var[z] ≈ 784标准差约 28。也就是说经过一层线性变换信号的规模被放大了近 30 倍。如果网络有 3 层、5 层每层都这样放大最后 logits 就是天文数字。这个现象用生活中的例子类比你往扬声器里说话如果每一级放大器都把信号放大 30 倍那么经过两级放大原本微弱的声音就变成震耳欲聋的噪音甚至烧坏设备。神经网络里的前向传播就是这样权重初始化太大相当于每一层都是一台过猛的放大器。2.2 反向传播中的梯度爆炸与梯度消失反向传播时误差信号从输出层往回传每一层都要乘上权重的转置矩阵和激活函数的导数。如果初始化权重值很大且激活函数导数在饱和区趋近 0这两股力量会把梯度推向两个极端靠近输出层的层梯度可能因为大权重连乘而爆炸靠近输入层的层梯度可能因为 sigmoid 饱和而消失。两种现象同时存在网络就彻底没法训练了。具体来说输出层的梯度dz2 probs - one_hot这个值通常不会特别大在 [−1,1] 范围内。但反向传播到W2时dW2 a1.T.dot(dz2)如果a1的数值很大或很稀疏dW2就会非常不稳定。更关键的是dz1 dz2.dot(W2.T) * sigmoid_derivative(z1)W2的数值如果很大dz1会爆炸而sigmoid_derivative(z1)在饱和区又几乎为 0。爆炸和消失交织梯度范数要么为inf要么为 0你根本无法判断该往哪个方向更新。我在调试时习惯打印每一层梯度的 L2 范数。当权重初始化过大时dW2的范数可能在 100 以上而dW1的范数却小于 1e-10。这种数量级的割裂就是大权重初始化引发的典型梯度失衡。2.3 激活函数饱和sigmoid 的死亡区域上面提到 sigmoid 的两个饱和区也就是输入绝对值大于 4 的时候输出几乎就是 0 或 1导数趋近 0。如果第一层加权和z1标准差达到 28那么绝大多数z1的绝对值远大于 4可以说整个隐藏层都陷入了死亡区域。有人可能会想那换 ReLU 呢ReLU 也有自己的问题如果权重初始化太大z1变成很大的负数ReLU 输出为 0反向传播梯度为 0神经元直接“坏死”而且再也无法恢复。所以无论用哪种激活函数大权重初始化的后果都是灾难性的只是表现略有不同sigmoid 是饱和不更新ReLU 是硬性死亡。理解这一点你就能明白为什么初始化尺度必须匹配激活函数的特性而不是随便选一个。3. 用 NumPy 复现大权重初始化问题3.1 MNIST 数据准备避开下载坑网上很多教程会用torchvision.datasets.MNIST来下载数据但由于托管地址经常变更我在写 NumPy 系列时遇到了 404很不稳定。所以这里直接教大家用本地文件读取的方式一次下载后续随便用。MNIST 官方数据是四个二进制文件分别是训练图像、训练标签、测试图像、测试标签。解压后可以通过struct.unpack读取文件头信息再解析像素数据。下面是我常用的加载函数import numpy as np import struct def load_mnist(path, kindtrain): labels_path f{path}/{kind}-labels.idx1-ubyte images_path f{path}/{kind}-images.idx3-ubyte with open(labels_path, rb) as lbpath: magic, n struct.unpack(II, lbpath.read(8)) labels np.fromfile(lbpath, dtypenp.uint8) with open(images_path, rb) as imgpath: magic, num, rows, cols struct.unpack(IIII, imgpath.read(16)) images np.fromfile(imgpath, dtypenp.uint8).reshape(num, rows * cols) return images.astype(np.float32) / 255.0, labels这段代码里II表示以大端方式读取两个无符号整数。MNIST 的文件格式就是大端存储这也是容易踩坑的地方——如果直接用np.loadtxt或者open乱读数据会完全错乱。读出来后要对像素值除以 255把数据范围归一化到 [0,1]。注意这只是最基础的归一化后面我们会看到[0,1] 的输入和高斯权重组合依然会产生大 z 值如果进一步做零均值标准化效果会更好但为了复现问题这里先保持 [0,1] 输入。3.2 两层神经网络的 NumPy 实现我写了一个极简的两层网络没有用任何框架方便你逐行追踪计算过程。激活函数选择 sigmoid损失函数用 softmax 交叉熵。def sigmoid(z): return 1 / (1 np.exp(-z)) def sigmoid_deriv(a): return a * (1 - a) def softmax(z): exp_z np.exp(z - np.max(z, axis1, keepdimsTrue)) return exp_z / np.sum(exp_z, axis1, keepdimsTrue) def cross_entropy_loss(y_pred, y_true): m y_true.shape[0] log_likelihood -np.log(y_pred[np.arange(m), y_true] 1e-8) return np.mean(log_likelihood)训练循环里每一步做前向传播、计算损失、反向传播、更新权重。这里使用批量梯度下降每次迭代用全部 60000 张图虽然慢但梯度方向更稳定便于观察数值现象。def train(X, y, W1, b1, W2, b2, lr, epochs): m X.shape[0] loss_history [] for epoch in range(epochs): # 前向传播 z1 X.dot(W1) b1 a1 sigmoid(z1) z2 a1.dot(W2) b2 probs softmax(z2) loss cross_entropy_loss(probs, y) # 反向传播 y_onehot np.zeros_like(probs) y_onehot[np.arange(m), y] 1 dz2 probs - y_onehot dW2 a1.T.dot(dz2) / m db2 np.mean(dz2, axis0) dz1 dz2.dot(W2.T) * sigmoid_deriv(a1) dW1 X.T.dot(dz1) / m db1 np.mean(dz1, axis0) # 更新 W2 - lr * dW2 b2 - lr * db2 W1 - lr * dW1 b1 - lr * db1 loss_history.append(loss) if epoch % 2 0: print(fepoch {epoch}, loss {loss:.4f}) return W1, b1, W2, b2, loss_history注意反向传播里dz1 dz2.dot(W2.T) * sigmoid_deriv(a1)这里必须使用当前z1对应的激活值计算导数。很多人会误写成sigmoid_deriv(z1)但sigmoid_deriv其实是a1 * (1 - a1)和传入z1的结果不同这一点特别容易写错建议亲手核对一下维度。3.3 大权重初始化实验randn 裸奔下面重点来了。我先用最原始的大权重初始化跑一次训练看看会发生什么。np.random.seed(42) W1 np.random.randn(784, 100) * 1.0 b1 np.zeros(100) W2 np.random.randn(100, 10) * 1.0 b2 np.zeros(10) W1, b1, W2, b2, loss_hist train(X_train, y_train, W1, b1, W2, b2, lr0.1, epochs10)运行结果如下我用文字记录关键数值第 1 个 epochloss 约 87.2345。这个值非常大因为 logits 数值巨大softmax 输出接近 one-hot交叉熵接近最大值ln(10)的很多倍不对这里是因为 logits 分布极端交叉熵可能出现很大值。第 3 个 epochloss 变成nan。此时z1的最小值和最大值大约为 -1478 和 1523你可以想象这个网络已经完全失控。为了看清楚问题出在哪一层我在训练循环里加入了中间变量统计print(z1 std:, np.std(z1)) print(a1 mean:, np.mean(a1)) print(dW1 norm:, np.linalg.norm(dW1)) print(dW2 norm:, np.linalg.norm(dW2))大权重下第一轮输出的z1标准差确实在 28 左右a1的均值接近 0.5但大部分激活值都集中在 0 和 1不是均匀分布。dW1的范数在 1e-5 级别而dW2的范数在 10 以上。这说明梯度传递在隐藏层出现了断裂。3.4 调整初始化到不同尺度观察梯度与损失的变化为了更深入对比我做了三组实验分别把权重缩放系数设为 1.0、0.1、0.01。scale 1.0第一层z1标准差约 28loss 迅速变为 NaN。scale 0.1z1标准差约 2.8sigmoid 饱和现象减轻但仍有大量接近 0 或 1loss 刚开始能降一点但后续震荡剧烈最终也无法收敛。scale 0.01z1标准差约 0.28sigmoid 工作在线性区其实 sigmoid 在线性区导数大梯度传递正常loss 平滑下降。由此你能直观看到一个边界当z1的标准差接近 1 时网络处于可训练状态超过 2-3就会出现明显的饱和与梯度失衡。这也为下面推导合理的初始化公式提供了经验佐证。4. 从数学推导到合理初始化4.1 全零初始化的隐藏陷阱大权重不行那全部初始化为 0 行不行也不行。试想一下如果W1全是 0那么所有隐藏神经元的输入都是 0经过 sigmoid 后输出都是 0.5且对于任何输入同一层所有神经元的行为完全一致。反向传播时dW1的每个元素都等于np.dot(X.T, dz1)而dz1的每一列相同所以dW1的每一行也是相同的。这意味着即使经过多轮梯度下降第一层权重的每个神经元依然保持着相同的变化模式永远无法分化出不同的特征。这就是对称性问题。所以初始化要在“太小”和“太大”之间取一个平衡既不能让信号逐层消失也不能让它逐层爆炸还要打破神经元之间的对称性。理解这个诉求后下面的方差保持法就很自然了。4.2 方差保持为什么 1/sqrt(n) 是关键回到前向传播的方差公式如果希望每一层输出的方差和输入方差大致相同就要让n * Var[w] 1也就是Var[w] 1/n即权重标准差为1/sqrt(n)。这里的n通常称为 fan_in也就是该层神经元接收的输入数量。对于输入层n784std 1/sqrt(784) ≈ 0.0357。我之前实验里 scale0.01 能正常工作0.1 就崩就是因为 0.01 非常接近 0.0357 这个量级而 0.1 大了接近三倍已经在方差放大线的边缘了。用这个公式初始化第一层z1的标准差约为 1信号进入激活函数时处于不饱和的敏感区间梯度可以顺畅回流。反向传播也有类似的方差约束。误差信号从输出层反向传播时经过权重矩阵W的转置其维度影响是上一层的输出维度fan_out。因此一种更对称的做法是取 fan_in 和 fan_out 的平均或调和关系这就是 Xavier 初始化的由来。4.3 Xavier 与 He 初始化到底在解决什么Xavier 初始化也叫 Glorot 初始化最早在 2010 年被提出它要求前向和反向传播时信息都能在其间流动方差保持一致。常见形式有两种一种近似为std sqrt(2 / (fan_in fan_out))另一种简单形式是std 1/sqrt(fan_in)。具体使用哪一种取决于你对“方差保持一致”的精确定义。对于 Sigmoid 激活Xavier 初始化在实践中表现良好因为 Sigmoid 在 0 附近近似线性线性网络的方差传播理论近似成立。He 初始化是专门为 ReLU 这类非饱和激活设计的。因为 ReLU 会把一半神经元置为 0实际传递的方差减半所以需要补偿一个sqrt(2)因子即std sqrt(2 / fan_in)你可以把 ReLU 想象成一个只允许正半轴信号通过的门负半轴的信息直接丢弃为了维持同样的统计强度初始权重就必须更大一点。简单记一个结论用 Sigmoid 或 Tanh优先选 Xavier用 ReLU 系列优先选 He如果拿不准就从std1/sqrt(fan_in)开始然后看激活值分布来微调。这个经验在实际调试中非常管用。5. 实操修复在 NumPy 中实现合理初始化5.1 三种初始化方式对比下面这段代码封装了三种不同初始化方式方便你切换实验。def init_weights_randn(fan_in, fan_out, scale1.0): return np.random.randn(fan_in, fan_out) * scale def init_weights_xavier(fan_in, fan_out): std np.sqrt(2.0 / (fan_in fan_out)) return np.random.randn(fan_in, fan_out) * std def init_weights_he(fan_in, fan_out): std np.sqrt(2.0 / fan_in) return np.random.randn(fan_in, fan_out) * std用 Xavier 初始化以后训练过程明显变得顺滑。为了验证效果我跑了 20 个 epoch学习率 0.1批量大小 60000。损失从初始的 2.3 左右逐渐下降到 0.35 附近测试准确率最终稳定在 97.2% 左右。相比之下大权重初始化版本在第一个 epoch 后就已经无法挽回了。5.2 训练结果对比与可视化思路虽然这里不能贴实时图表但你可以自己把loss_history画出来。大权重初始化的 loss 曲线是一条直线冲上天的样子前两个点还是几十第三个点直接nan合适初始化的 loss 曲线则是一条平稳下降的曲线没有毛刺没有跳变。如果你要输出准确率可以加一段预测函数def predict(X, W1, b1, W2, b2): z1 X.dot(W1) b1 a1 sigmoid(z1) z2 a1.dot(W2) b2 return np.argmax(softmax(z2), axis1) test_acc np.mean(predict(X_test, W1, b1, W2, b2) y_test)我实测下来Xavier 初始化配合 sigmoid在两层网络上测试准确率大约在 97%具体数值会随随机种子和超参数小幅波动。如果换用 He 初始化 ReLU相同结构在 MNIST 上甚至能到 97.8% 左右。这也印证了初始化要和激活函数匹配。5.3 一个更稳妥的调试流程我建议你在每次训练之前先做一次“初始化冒烟测试”随机初始化后取一小批数据跑一次前向传播打印每一层的激活值分布均值、方差、最小最大值。如果发现激活值集中在 0 或者 1说明初始化尺度太大如果激活值集中在 0.5 附近但方差极小可能是初始化尺度太小梯度会消失网络学不动。类似地反向传播时打印dW1、dW2的范数。一个健康的网络各层梯度的范数应该在同一个数量级附近不会差出 10 个数量级。我见过不少人用框架训练时隐藏层梯度范数在 1e-4输出层在 1e-2虽然能训练但速度极慢问题多半也出在初始化上。6. 常见问题与排查技巧实录6.1 训练一开始 loss 就是 NaN怎么办按顺序排查第一看输入数据是否归一化。如果你的像素值还是 0-255 的整数输入方差非常大大权重初始化会放大得更厉害。第二看权重初始化尺度。检查std np.std(W1)如果大于1/sqrt(784)好几倍那就先把初始化改合理。第三看学习率。如果前两者都正常学习率设为 0.1 对两层网络是合适的但对更深网络可能偏大需要适当减小。第四看交叉熵实现有没有做数值稳定比如 softmax 里先减去最大值。我遇到很多次 NaN最后发现是交叉熵里np.log(0)导致-inf而-inf在反向传播里变成 NaN。加一个小的 epsilon 比如1e-8能临时避免但根源还是probs里面出现了 0这通常又是 logits 太大造成的所以还是回到初始化。6.2 如何确认反向传播写对了梯度检查如果你手写网络反向传播的细微错误会导致初始化问题和非初始化问题混在一起所以强烈建议做一次梯度检查。数值梯度公式很简单def compute_gradient_numerical(X, y, W, epsilon1e-6): grad np.zeros_like(W) it np.nditer(W, flags[multi_index]) while not it.finished: idx it.multi_index old W[idx] W[idx] old epsilon loss_plus cross_entropy_loss(predict_softmax(X, W), y) W[idx] old - epsilon loss_minus cross_entropy_loss(predict_softmax(X, W), y) W[idx] old grad[idx] (loss_plus - loss_minus) / (2 * epsilon) it.iternext() return grad然后对比解析梯度和数值梯度相对误差应该在 1e-6 以下。如果误差很大先修反向传播再谈初始化。这个步骤很多人跳过了结果代码里的 bug 和初始化问题混在一起调参调半天也没用。6.3 激活函数选型对初始化的影响Sigmoid 在 MNIST 这种浅层网络上基本够用但它确实有饱和问题。如果换了 Tanh它的输出范围是 [-1,1]比 Sigmoid 好一些但仍然存在饱和区。Xavier 初始化对 Tanh 更友好。如果换 ReLU就改用 He 初始化训练速度会明显加快。我在实验中发现即使是合适的 Xavier 初始化Sigmoid 网络在 MNIST 上也只能达到约 97% 的准确率而 ReLU He 初始化相同结构能到 97.8%。对于这个简单任务差别不算大但对更深的网络激活函数和初始化的搭配是生死攸关的。6.4 手写网络和框架训练行为不一致有时你在 NumPy 里复现的模型和 PyTorch / TensorFlow 里的同结构训练结果差异巨大。框架默认的初始化策略通常是经过调优的比如 PyTorch 的nn.Linear默认使用 Kaiming Uniform 初始化和你的randn初始化不是一个分布。所以不要拿框架结果直接对比你的手写代码除非你手动把框架初始化也改成相同的策略。另外框架的自动求导使用的是 double 精度还是 float 精度也会影响数值行为。NumPy 默认浮点数是 float64而 PyTorch 默认是 float32虽然 MNIST 这种简单任务不会造成太大差异但在极端初始化下float32 更容易溢出出现 NaN。我建议手写复现时可以在初始化后加一句astype(np.float32)模拟真实框架环境更容易复现出框架里的 NaN 问题。几句实在话踩过大权重初始化的坑之后我养成了一个习惯不管用什么框架开始训练前先跑一小批数据打印每层激活值的直方图和梯度范数。这个过程不到一分钟能筛掉 80% 的数值不稳定问题。如果你也正在手写神经网络建议你亲手把scale从 3.0 到 0.001 都试一遍用实验体会“方差传导”这件事比看一百遍公式都深刻。这个系列后面还会拆解学习率、正则化、批量归一化的底层计算但作为地基先把初始化这块啃透后面你会走得轻松很多。
返回列表