如何在PyTorch中为自定义模型设置Xavier初始化权重?
在PyTorch中为自定义模型设置Xavier权重初始化
方法一:手动为每个线性层初始化
直接在模型的__init__方法中,对每个nn.Linear层的权重和偏置应用Xavier初始化。由于你使用sigmoid作为激活函数,推荐用Xavier均匀初始化(xavier_uniform_),它针对饱和型激活函数做了适配。
修改后的完整模型代码:
import torch import torch.nn as nn class Model(nn.Module): def __init__(self): super(Model, self).__init__() self.linear1 = nn.Linear(2, 512*8) self.linear2 = nn.Linear(256*16, 256*8) self.linear3 = nn.Linear(256*8, 1) # 对每层应用Xavier初始化 nn.init.xavier_uniform_(self.linear1.weight) nn.init.zeros_(self.linear1.bias) # 偏置通常初始化为0 nn.init.xavier_uniform_(self.linear2.weight) nn.init.zeros_(self.linear2.bias) nn.init.xavier_uniform_(self.linear3.weight) nn.init.zeros_(self.linear3.bias) def forward(self, x): x = self.linear1(x) x = torch.sigmoid(x) x = self.linear2(x) x = torch.sigmoid(x) x = self.linear3(x) return x
方法二:自动遍历所有线性层初始化
如果模型包含大量线性层,手动逐个初始化效率低,可以写一个初始化函数,通过apply()方法自动遍历并应用到所有线性层:
import torch import torch.nn as nn class Model(nn.Module): def __init__(self): super(Model, self).__init__() self.linear1 = nn.Linear(2, 512*8) self.linear2 = nn.Linear(256*16, 256*8) self.linear3 = nn.Linear(256*8, 1) # 自动初始化所有线性层 self.apply(self._init_weights) def _init_weights(self, module): # 仅对线性层执行初始化逻辑 if isinstance(module, nn.Linear): nn.init.xavier_uniform_(module.weight) if module.bias is not None: nn.init.zeros_(module.bias) def forward(self, x): x = self.linear1(x) x = torch.sigmoid(x) x = self.linear2(x) x = torch.sigmoid(x) x = self.linear3(x) return x
补充说明
- Xavier初始化的核心是让每层输入和输出的方差尽可能一致,避免梯度消失或爆炸,非常适合搭配
sigmoid、tanh这类饱和型激活函数。 - 如果后续改用ReLU类激活函数,更推荐使用He初始化(
nn.init.kaiming_uniform_)。
内容的提问来源于stack exchange,提问作者asdfe
相关产品推荐
相关产品推荐

