You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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()
      
  • 方法二:修改模型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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.25 09:45:34