ARTICLE DETAIL

资讯详情

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

PyTorch参数初始化实战:从梯度消失到Kaiming初始化的原理与排查

PyTorch参数初始化实战:从梯度消失到Kaiming初始化的原理与排查 今天聊PyTorch里的参数初始化这是神经网络基础中一个特别容易被跳过的知识点但恰恰是它藏着训练能不能顺利起步的秘密。我最早做模型时网络搭完直接丢给优化器完全不关心默认初始化是什么结果训练loss跟心电图一样乱跳折腾好几天才发现问题根本不在数据而在于初始权重分布。这篇文章既是对这个问题的复盘也把这些年用到的初始化知识系统梳理一遍适合正在啃PyTorch框架和神经网络基础的朋友也适合已经跑过一堆实验但对“setting seed”只知其然的调参党。看完你至少能搞清楚四件事为什么初始化这么重要、主流方法到底在解决什么、代码里能怎么落地、以及初始化不对时该怎么排查。1. 初始化的分量为什么它是训练开始前最关键的“隐形参数”1.1 从梯度消失与梯度爆炸说起初始化如何左右整个训练进程我习惯把参数初始化称作训练开始前的“隐形变量”。你会调学习率、会换优化器、会加正则项但你可能从来没有去修改过模型里那一堆weight和bias是怎么来的。实际上一次训练收敛得好不好、稳不稳、快不快从网络参数被创建出来的那一刻就已经被烙下了印记。先理一下最底层的关系。神经网络训练说到底就是梯度下降而梯度下降需要一个起点。起点不一样后面走的路就完全不一样。你可以想象一个山地场景训练的目标是找到一个低洼地也就是损失函数的最小值而参数初始位置决定你从哪个点开始往下走。如果起点放在一片非常平缓的高原梯度近似为0那你无论怎么走都寸步难行如果起点放在一个非常陡峭的悬崖边缘稍微迈一小步就会翻滚出去loss直接飞向NaN。但仅仅用“起点”来解释还不够。参数初始化真正厉害的地方在于它通过控制初始权重的方差直接影响网络在训练初期的信号传播进而决定会不会触发梯度消失和梯度爆炸。这里面的机制值得展开说。假设我们有一个L层全连接网络先忽略激活函数每一层都是一个线性变换y Wx。如果W中每个元素都是从均值为0、标准差为σ的分布中独立采样那么经过一层线性变换后输出y的方差近似等于Var(y) ≈ fan_in * σ² * Var(x)其中fan_in是输入维度。这个公式怎么来的简单来说y的每一项是输入向量x与权重矩阵一行的内积。如果x的各分量独立且同分布权重也是独立同分布那么连加后的方差自然就是n个乘积方差之和。这也就是方差传递的核心关系。把这个关系连续叠L层会变成什么样Var(y_L) ≈ (fan_in * σ²)^L * Var(x)。如果fan_in * σ² 1方差逐层指数放大到第20层时可能从数字1变成天文数字再往后就是NaN如果fan_in * σ² 1方差逐层指数缩小到第20层时输出几乎全部挤到0附近梯度消失。所以你会发现初始化理论的本质就是一件事让每一层输入输出的方差尽量保持一致既不放大也不衰减。后面要讲的Xavier和Kaiming全都是围绕这个“方差守恒”目标展开的。1.2 先弄懂PyTorch里的参数长什么样Parameter、Tensor与默认初始化既然要聊PyTorch里的参数初始化就得先看清楚参数在框架里到底以什么形式存在。在PyTorch中一个模型的权重不是普通的Tensor而是nn.Parameter它是一个特殊的Tensor子类在模块中注册为parameter。只要模块被注册后优化器就能通过model.parameters()收集到它并自动参与梯度更新。当你执行model nn.Sequential(...)或者自定义一个子类继承nn.Module时层内部的weight和bias已经用默认策略初始化过了。这一点非常重要即使你什么都没做PyTorch也已经给了你一个“默认初始化”。比如nn.Linear的weight默认用kaiming_uniform_初始化bias默认按均匀分布初始化。很多人会在无意中依赖这些默认行为但完全不清楚默认背后是什么。这也解释了为什么有些时候同一个模型不同人跑结果不同——因为默认初始化依赖全局随机种子而很多人并没有固定种子。如果你去看PyTorch源码Linear的reset_parameters方法是这样的def reset_parameters(self) - None: init.kaiming_uniform_(self.weight, amath.sqrt(5)) if self.bias is not None: fan_in, _ init._calculate_fan_in_and_fan_out(self.weight) bound 1 / math.sqrt(fan_in) if fan_in 0 else 0 init.uniform_(self.bias, -bound, bound)注意这里kaiming_uniform_传入了amath.sqrt(5)这个a是LeakyReLU的负斜率参数说明PyTorch默认假设你用的是ReLU族激活函数并且默认负斜率设为sqrt(5)。很多人从来没在意过这个细节但它直接影响了初始分布的形状。记住这个默认行为后面排查问题时你会回来感谢它的。2. 主流初始化方法原理拆解从零初始化到Kaiming/He初始化2.1 零初始化与常数初始化的致命缺陷对称性陷阱先来看最直观也是最危险的两种零初始化与常数初始化。所谓零初始化就是把所有权重设成0。权重为0意味着什么呢前向传播时不管输入是什么神经元的线性输出全为0或者统一等于一个常数偏置激活函数一作用所有神经元输出都相同。反向传播时因为损失函数对每个权重的偏导在权重等于0时也会呈现某种对称性所有参数会以完全相同的方式更新。结果是不论你的网络有多少个神经元实际上只有一个“有效神经元”在工作其他神经元都在做完全一样的事情。这一现象叫对称性陷阱。有人可能会想我把权重设成同一个常数比如0.5总行了吧很遗憾效果跟零初始化一样。只要所有权重相等无论常数是多少所有神经元的前向输出、反向梯度都相同对称性还是打不破。网络再宽、再深也只是一个线性退化模型。那权重设置成很小的数行不行比如normal_(mean0, std0.01)这得看网络深度。用一个很浅的网络std0.01确实能训练起来但是在深层网络中前面推导过的方差连乘效应马上就会显现0.01的初始方差经过20层之后输出方差变成原来的0.01^20级别基本就是0。这时梯度也跟随激活值一起消失模型就跟冻住了一样。2.2 随机初始化的方差控制前向传播与反向传播的守恒条件所以正确的思路不是放弃随机而是让随机有度。PyTorch中torch.nn.init提供了normal_和uniform_两种基础随机方案但直接用它们的问题在于你必须自己指定方差。若方差设得太大经过多层后输出爆炸设得太小经过多层后输出消失。这就是前面说的方差守恒条件。这里要引入一个概念fan_in和fan_out。对全连接层来说fan_in是输入神经元数量fan_out是输出神经元数量对卷积层来说fan_in不再只是通道数而是“输入通道数 × 卷积核大小”fan_out则是“输出通道数 × 卷积核大小”。PyTorch在计算初始化bound时用的就是这两个值。理解fan_in和fan_out的意义在于它们直接决定了权重方差应该取多大。前向传播需要保留的是输入信号的方差所以跟fan_in相关反向传播需要保留的是梯度信号的方差所以跟fan_out相关。一个初始化方法是否优秀核心就看它在多大程度上同时满足这两边的需求。2.3 Xavier/Glorot初始化面向tanh与sigmoid的均衡方案在这个背景下Xavier初始化也就是Glorot初始化登场了。它解决的是使用tanh、sigmoid这类饱和激活函数时如何让前向与反向传播的方差都保持稳定的问题。推导逻辑大致是为了让前向传播每层输出方差不变要求权重方差约等于1/fan_in为了让反向传播每层梯度方差不变要求权重方差约等于1/fan_out。前向和反向的要求没法同时满足除非fan_in等于fan_out。所以Xavier取了一个调和平均让权重方差落在2/(fan_in fan_out)附近。落到均匀分布时上界就是那个经典形式U(-sqrt(6/(fan_infan_out)), sqrt(6/(fan_infan_out)))。PyTorch里xavier_uniform_和xavier_normal_干的就是这件事。它们还有一个gain参数用来配合激活函数做补偿tanh的gain默认是5/3线性激活是1。但在实际使用中我发现很多人拿Xavier初始化去配ReLU效果其实一般。原因在于Xavier假设激活函数是线性的或者关于原点对称的而ReLU无脑砍掉一半输出导致方差直接减半。这正好引出下一个方法。2.4 Kaiming/He初始化ReLU时代的事实标准Kaiming初始化也常叫He初始化正是针对ReLU及其变体设计的。它的关键洞察是ReLU会把小于0的输入全部置0导致实际输出方差大约只有输入方差的一半。为了补偿这个“能量损失”Kaiming初始化把权重方差从Xavier的均衡方案改成了2/fan_in。分母不是fan_infan_out而是只有fan_in原因是它更优先保证前向传播当然也有研究者按fan_out计算对应mode的不同取值。PyTorch里kaiming_normal_和kaiming_uniform_就是这套逻辑的实现。modefan_in保留前向传播信号modefan_out保留反向传播梯度。对于卷积层而言fan_in和fan_out会按照感受野范围自动计算而不是简单按输入输出通道数来算。这一点对于理解卷积网络的初始化特别重要。这一设计直接帮助了后来一系列深层ReLU网络的训练包括VGG、ResNet。即便今天你已经不怎么手动初始化了PyTorch里大多数卷积层和全连接层默认用的也是Kaiming系初始化。所以你会发现一个现代CNN从零训练其实就是默认在Kaiming初始化上做梯度下降。新手阶段不需要魔改但必须知道默认值是什么、为什么要这么设。3. PyTorch中参数初始化的几种落地姿势3.1 直接调用torch.nn.init最直白的路子PyTorch的torch.nn.init模块提供了非常完整的初始化函数。最直白的用法就是对某个层的参数单独操作import torch import torch.nn as nn import torch.nn.init as init model nn.Sequential( nn.Linear(256, 128), nn.ReLU(), nn.Linear(128, 10) ) # 对某一层单独做初始化 init.xavier_normal_(model[0].weight, gain1.0) init.kaiming_normal_(model[2].weight, modefan_in, nonlinearityrelu) init.constant_(model[0].bias, 0.0)注意这些init函数都是原地操作也就是in-place修改传入的Tensor同时也会返回同一个Tensor所以不需要重新赋值。这一点跟很多函数式API不一样初学容易误以为init.xavier_normal_返回了新参数而忘了原参数早被改了。另外要留心的是这些函数只作用于单一的Tensor。如果你的模型有二三十个层一个个手动初始化显然不现实。这就是第二种姿势登场的地方。3.2 自定义初始化函数配合model.apply()灵活且干净PyTorch的nn.Module有一个apply方法它会递归地对模块下所有子模块执行传入的函数。这是做全模型初始化的标准姿势def init_weights(m): if isinstance(m, nn.Linear): init.kaiming_normal_(m.weight, modefan_in, nonlinearityrelu) if m.bias is not None: init.zeros_(m.bias) elif isinstance(m, nn.Conv2d): init.kaiming_normal_(m.weight, modefan_in, nonlinearityrelu) if m.bias is not None: init.zeros_(m.bias) elif isinstance(m, nn.BatchNorm2d): init.ones_(m.weight) init.zeros_(m.bias) model.apply(init_weights)这个写法有两点值得强调。第一因为apply是递归遍历的init_weights会被每个子模块调用所以你一定要用isinstance判断当前模块是什么类型否则有可能对Sequential、ModuleList这种容器模块也执行操作或者漏掉某些自定义层。第二很多自定义模块里内嵌了多个Linear或Convisinstance判断能自然覆盖到它们因为apply遍历到最底层时每个Linear都会被单独调用一次init_weights。我自己的习惯是将这类初始化函数单独放在一个utils模块里原因很简单训练过程中如果要做对比实验比如验证“标准差为0.1的正态初始化”和“Kaiming初始化”哪个更适合当前任务只要切换传入apply的函数即可完全不改动模型代码。3.3 从源码层面重写reset_parameters()理解默认初始化的真相更进阶的姿势是重写模块的reset_parameters方法。当你自定义一个层时PyTorch并不会给你自动初始化参数除非你在__init__里手动调用reset_parameters。比如import math class MyLinear(nn.Module): def __init__(self, in_features, out_features): super().__init__() self.weight nn.Parameter(torch.empty(in_features, out_features)) self.bias nn.Parameter(torch.empty(out_features)) self.reset_parameters() def reset_parameters(self): init.kaiming_uniform_(self.weight, amath.sqrt(5)) fan_in, _ init._calculate_fan_in_and_fan_out(self.weight) bound 1 / math.sqrt(fan_in) if fan_in 0 else 0 init.uniform_(self.bias, -bound, bound)这里用到了init._calculate_fan_in_and_fan_out这是一个带下划线的“私有”函数但在PyTorch源码中它自己也是这么用的实际项目中完全可以直接调用。重写reset_parameters的好处是你的自定义层从创建那一刻就带着一套明确的初始化策略后续调用model.apply(my_custom_layer.reset_parameters)时也能精准命中不会误伤其他层。我在实践中发现初学者对自定义层的初始化往往有两条误区。一是以为继承了nn.Module就会自动初始化所有参数结果自定义的Parameter全是torch.empty里的随机垃圾值训练直接不收敛。二是以为bias可以不初始化实际上bias如果保留初始的随机值且量级比权重大前向输出会被bias主导激活函数很容易进入饱和区。无论你用什么初始化策略bias初始化为0或很小的常数通常是安全选择。4. 初始化不当的典型症状与一套可复现的排查链路4.1 三个典型症状NaN、loss纹丝不动、梯度值异常先说最常见的三种“病”都是我这些年实际见过的每种都有明确的指向性。第一种是训练一开始loss就直接变成NaN。这种情况往往是初始权重标准差过大配合上比较大的学习率数值在深层前向传播中溢出。也可能是深层网络的激活值太大经过几层后直接让loss计算出现inf然后反向传播时梯度也跟着inf最终整个参数变成NaN。如果数据本身没问题第一个要怀疑的对象就是初始化方差。第二种是loss纹丝不动或者说下降极慢像一条平线。这种情况最大的嫌疑是权重初始化方差过小梯度消失模型根本没有在学。你可能会看到loss在某个值附近轻微抖动但无论怎么调学习率都不改善。如果你用的是RNN/LSTM还需要额外怀疑时间维度上的梯度消失。但如果是普通前馈网络初始化方差过小是最常见的原因。第三种是loss能下降但震荡非常剧烈训练曲线像锯齿一样。这往往说明初始权重绝对值过大模型一开始就走到了损失曲面非常陡峭的位置各个batch之间的梯度方向极不稳定。这种症状虽然没有前面两种致命但会拖慢收敛而且很容易让你误判学习率是主要问题。4.2 逐层诊断法用钩子函数定位出问题的层如果你只是怀疑初始化有问题但不知道具体是哪一层出的问题那就需要一套可复现的排查链路。我的做法是这样。第一步固定随机种子。这一步不做后面所有对比都不可信。固定种子的写法很简单import torch import numpy as np import random def set_seed(seed42): torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) np.random.seed(seed) random.seed(seed)第二步检查输入数据的分布。数据如果没做标准化输入数值本身就有几百上千的量级那再好的初始化也扛不住。确认输入均值和方差都在合理范围后再来看模型内部。第三步用forward_hook逐层打印激活值的均值和标准差。这是定位问题层的核心手段activations {} def make_hook(name): def hook(module, input, output): if isinstance(output, torch.Tensor): activations[name] output.detach() return hook for name, module in model.named_modules(): if isinstance(module, (nn.Linear, nn.Conv2d)): module.register_forward_hook(make_hook(name)) # 喂一个批次的数据 model(x) for name, act in activations.items(): print(f{name}: mean{act.mean():.4f}, std{act.std():.4f})看这些std的变化趋势如果从第一层到最后一层std在逐层指数级变大那就是初始化方差偏大如果逐层指数级变小那就是初始化方差偏小。通常到第10层左右std就会急剧偏离正常范围问题层的位置一目了然。第四步检查反向梯度。你也可以用register_full_backward_hook来抓取不同层的梯度均值如果梯度从最后一层往前逐层缩小很快比如到第5层梯度已经变成1e-12级别那么梯度消失的方向基本确定了。4.3 修复示例替换初始化前后训练曲线发生了什么定位到问题层之后修复起来其实很直接。比如你发现某个Transformer编码块里的全连接层在第6层之后std突然掉到1e-6那你就可以单独对该层做一次Kaiming初始化再跑一次forward看输出分布是否恢复正常。下面是一个我实际用过的例子。假设模型是一个20层的纯全连接网络我用std0.01初始化训练200步之后loss几乎不下降。直接换成kaiming_normal_后同样是200步loss下降曲线立刻有了明显斜率。差别就是这么直观。这里的关键结论是初始化不对时你怎么调学习率、怎么换优化器都是徒劳因为梯度信号本身就已经被初始分布抑制了。替换初始化时还有几个注意点初始化只应该在训练循环开始前执行一次不要在每次epoch之后重复apply如果用了预训练权重千万不要盲目apply自定义初始化否则辛辛苦苦预训练出来的权重直接被覆盖另外如果模型里有BatchNorm普通情况下不需要额外初始化BN的weight和bias保留weight1、bias0即可乱改BN参数反而会影响训练稳定性。5. 参数初始化的选型建议与个人经验沉淀5.1 不同网络结构下的初始化策略对照表我把平时最常用的初始化策略整理成了一张对照表基本覆盖了主流网络结构网络结构/层类型推荐初始化备选方案说明全连接层 ReLUkaiming_normal_(fan_in)kaiming_uniform_PyTorch大部分默认即此策略全连接层 tanh/sigmoidxavier_normal_xavier_uniform_需要按激活函数调整gain卷积层 ReLUkaiming_normal_(fan_in)kaiming_uniform_卷积下fan_in自动算感受野LSTM/GRU隐层orthogonal_uniform_(-1/sqrt(hidden_size), 1/sqrt(hidden_size))PyTorch内置LSTM即用正交初始化Embedding层normal_(mean0, std1)xavier_normal_可用Word2Vec等预训练向量替代BatchNorm层保留默认不推荐修改weight1, bias0新加的迁移学习分类头kaiming_normal_ 小学习率normal_(std0.01)和冻结层配合使用这里单独说说LSTM。PyTorch的nn.LSTM在源码中会使用orthogonal_初始化部分权重尤其是起到跨时间步传播作用的hh权重。原因在于正交矩阵能保持向量范数在时间维度展开时相当于一系列正交矩阵连乘矩阵的特征值都落在1附近梯度在时间维度上就不容易指数级消失或爆炸。这一点对于训练长序列尤其重要。5.2 迁移学习、批归一化与残差连接场景下的新课题当网络结构里出现BatchNorm、LayerNorm和残差连接时初始化的风险会被架构天然缓解不少。我发现很多初学者会有一种错觉既然BN能强行把输出拉回标准正态那初始化是不是随便设一下就行这个想法有一半对也有一半错。对的那一半是BN确实让前向方差不再容易爆炸它等于在每个层之后加了一道“归一化闸门”。错的那一半是初始化绝对值的大小依然会影响训练初期BN里running_mean和running_var的更新速度以及反向传播时梯度的量级。更关键的是如果网络深到一定程度没有残差连接的情况下单纯靠Kaiming初始化也无法保证几十层ReLU网络的方差始终稳定所以现代深度网络普遍依赖残差连接和归一化来兜底。迁移学习场景下初始化策略又不一样。骨干网络通常加载预训练权重因此不会重新初始化。而你新加的分类头我建议不要用太大的初始化方差。虽然Kaiming初始化在新网络里很好用但在迁移学习里分类头的初始输出直接决定了模型第一次前向时的loss量级如果输出过大一开始的loss会比期望高出几个数量级带动梯度剧烈震荡。经验做法是用默认Kaiming或者稍小的正态分布同时把分类头的学习率调成骨干网络的十分之一左右。5.3 我踩过的初始化相关的坑与最终养成的习惯最后分享几个我自己的实操习惯都是踩过坑之后换来的。第一个习惯是任何新模型开跑之前先跑一个batch的数据把forward的激活值std粗略打印一遍。这几乎不花时间却能提前发现很多隐患。我遇到过不止一次模型结构完全没问题、数据也没问题但loss就是起飞检查激活值分布才发现某个自定义层的权重根本没人初始化过里面全是垃圾值。第二个习惯是所有对比实验固定同一个随机种子。以前做实验时同一套代码跑两次结果差了0.5个点我还以为是优化器的问题后来发现纯粹是没固定种子初始化的随机性掩盖了真实差异。现在我把set_seed(42)写进了统一训练模板里对比实验的可靠性高了很多。第三个习惯是RNN/LSTM网络如果不收敛先检查是不是初始化问题。时间步展开让梯度问题被放大十倍而orthogonal初始化往往能解决一大半。全连接和卷积网络还可以慢慢排查RNN/LSTM这种时间维度上的结构哪怕用了好的初始化也建议配合梯度裁剪一起用效果会稳定很多。最后分享一个小技巧。当你实在不知道该用什么初始化时可以先跑一小段时间的warmup训练比如20步然后打印各层权重的标准差变化。如果一个层的权重标准差在训练初期剧烈膨胀说明初始方差给大了如果几乎不动说明初始方差给小了。这个技巧虽然土但比任何纸上谈兵都实用尤其是在你面对一个全新的自定义网络结构时。
返回列表