PyTorch中nn.Linear默认使用均匀分布U(−1/√in_features, 1/√in_features)初始化权重,但该策略对ReLU等激活函数不友好,易致梯度消失或爆炸;正确做法是依激活函数选Kaiming(ReLU)或Xavier(Sigmoid/Tanh)初始化,并显式调用nn.init相应函数初始化weight和bias。

Linear层权重初始化的常见错误和正确做法
直接用 torch.nn.Linear 构造后不显式初始化,权重默认是均匀分布(U(-1/sqrt(in_features), 1/sqrt(in_features))),但这个初始范围对深层网络或特定激活函数(如ReLU)并不友好,容易导致梯度消失或爆炸。
推荐按激活函数类型选择初始化方式:
- ReLU 或其变体(LeakyReLU、PReLU):用
torch.nn.init.kaiming_normal_,设nonlinearity='relu',并确保mode='fan_in'(默认) - Sigmoid / Tanh:用
torch.nn.init.xavier_normal_ - 不做任何非线性变换(如最后输出层):可考虑
torch.nn.init.xavier_uniform_或手动截断正态分布
示例:
linear = torch.nn.Linear(128, 64) torch.nn.init.kaiming_normal_(linear.weight, nonlinearity='relu') torch.nn.init.zeros_(linear.bias) # 偏置通常清零
Conv2d层权重初始化时容易忽略的维度顺序
torch.nn.Conv2d 的权重张量形状是 (out_channels, in_channels, kernel_size[0], kernel_size[1]),而 torch.nn.init 系列函数默认按第一个维度(即 out_channels)计算 fan-in / fan-out —— 这对卷积层是合理的,但必须确认你没把 ConvTranspose2d 或分组卷积混进来。
立即学习“Python免费学习笔记(深入)”;
关键点:
图片提示词生成器?不止如此。 马甲系统 —— 把脑海中的画面,翻译成AI能理解的专业表达。 用得越多,它越懂你:首次需要多问几句确认方向,用久了几乎一说就懂。 用得越多,它越快:缓存机制让后续对话越来越省。 RAG进化:成功案例持续入库,越跑越聪明。 输入「新手指南」查看完整功能介绍
-
kaiming_normal_和xavier_normal_都能直接作用于conv.weight,无需 reshape - 如果用了分组卷积(
groups > 1),fan-in 自动按每组输入通道数计算,一般不用干预 - 切勿对
conv.weight.data手动调用torch.randn后不缩放 —— 这会绕过 fan 计算逻辑,破坏初始化理论依据
错误示范:conv.weight.data = torch.randn(conv.weight.shape) → 权重标准差为1,远大于合理初始值。
如何统一初始化整个模型的所有Linear/Conv层
逐层手动初始化易遗漏,尤其模型嵌套深时。更可靠的方式是在模型定义后遍历 modules(),按类型分发初始化逻辑:
def init_weights(m):
if isinstance(m, torch.nn.Linear):
torch.nn.init.kaiming_normal_(m.weight, nonlinearity='relu')
if m.bias is not None:
torch.nn.init.zeros_(m.bias)
elif isinstance(m, torch.nn.Conv2d):
torch.nn.init.kaiming_normal_(m.weight, nonlinearity='relu')
if m.bias is not None:
torch.nn.init.zeros_(m.bias)
<p>model.apply(init_weights)注意:apply() 会递归进入所有子模块,包括 Sequential、ModuleList 内部;但不会处理被注册为普通属性(而非 nn.Module 子类)的层,比如手写的 self.my_layer = nn.Linear(...) 是有效的,而 self.my_layer = Linear(...)(没加 nn.)则不会被识别。
BatchNorm层的weight/bias要不要初始化?
要,但方式不同:BatchNorm2d 的 weight(即 gamma)默认初始化为1,bias(即 beta)默认为0 —— 这已符合常用实践,一般无需改动。但如果你在残差连接后接 BN,有时会把 gamma 初始化为0(即「zero-init」),让初始状态等价于恒等映射,这对训练深层 ResNet 类模型有帮助。
操作方式:
- 保持默认:什么都不做
- zero-init BN scale:
torch.nn.init.zeros_(bn.weight)(注意不是bn.gamma,而是bn.weight) - 慎用
torch.nn.init.constant_(bn.bias, 0.1)—— 除非你明确知道偏置偏移对下游的影响
BN 层的初始化影响的是前向传播的缩放平移行为,和 Linear/Conv 的 fan-based 初始化逻辑完全不同,混用会导致收敛异常。

















