PyTorch中如何仅禁用nn.Dropout而不使用model.eval()?
仅禁用Dropout而保留BatchNorm正常工作的几种方法
方法一:手动遍历模块单独设置状态
无需修改模型结构,生成数据前可手动调整Dropout和BatchNorm的训练状态:- 若希望生成时BatchNorm使用当前batch的均值方差(保持训练模式):
# 先让模型整体处于训练模式 generator.train() # 遍历所有模块,将Dropout设为eval模式(禁用丢弃) for m in generator.modules(): if isinstance(m, nn.Dropout): m.eval() - 若希望生成时BatchNorm使用训练阶段积累的均值方差(eval模式):
# 先让模型整体处于eval模式,此时Dropout已禁用 generator.eval() # 遍历所有模块,将BatchNorm重新设为训练模式 for m in generator.modules(): if isinstance(m, nn.BatchNorm1d): m.train()
- 若希望生成时BatchNorm使用当前batch的均值方差(保持训练模式):
方法二:修改模型forward方法,添加开关控制Dropout
这种方式灵活性更高,直接在前向传播逻辑中加入参数控制Dropout是否生效,完全不受模型整体train/eval模式影响:
修改后的Generator代码:class Generator(nn.Module): def __init__(self, num_input=2, noise_dim=1, num_output=5, hidden_size=128): super(Generator, self).__init__() self.fc_in = nn.Linear(num_input+noise_dim, hidden_size) self.fc_mid = nn.Linear(hidden_size+num_input+noise_dim, hidden_size) self.fc_out = nn.Linear(2*hidden_size+num_input+noise_dim, num_output) self.bn_in = nn.BatchNorm1d(hidden_size) self.bn_mid = nn.BatchNorm1d(hidden_size) self.dropout = nn.Dropout() self.relu = nn.ReLU() def forward(self, y, z, use_dropout=True): h0 = torch.concat([y,z],axis=1) h1 = self.relu(self.bn_in(self.fc_in(h0))) # 仅当use_dropout为True时执行Dropout if use_dropout: h1 = self.dropout(h1) h1 = torch.concat([h0,h1],axis=1) h2 = self.relu(self.bn_mid(self.fc_mid(h1))) if use_dropout: h2 = self.dropout(h2) h2 = torch.concat([h1,h2],axis=1) x = self.fc_out(h2) return x训练时正常调用
generator(y, z)(默认启用Dropout),生成数据时调用generator(y, z, use_dropout=False)即可禁用Dropout,BatchNorm的行为可根据需求通过模型的train/eval模式自由设置。
内容的提问来源于stack exchange,提问作者Tota
相关产品推荐
相关产品推荐

