【问题标题】:What does model.train() do in PyTorch?model.train() 在 PyTorch 中做了什么?
【发布时间】:2018-12-28 05:23:56
【问题描述】:

它是否在nn.Module 中调用forward()?我想当我们调用模型时,正在使用forward 方法。 为什么我们需要指定 train()?

【问题讨论】:

  • 这些天在 PyTorch 中存在一个文档:pytorch.org/docs/stable/generated/torch.nn.Module.html 你可以查看文档,我认为它描述得很清楚。其他库/框架可能缺少文档,但在 PyTorch 中我认为官方文档非常好。
  • 或许“configure_training”或“set_training_mode”更适合这个函数。
  • 它通过 self.training = training 简单地通过 self.train(False) 为所有模块递归地更改 self.training。事实上,self.train 所做的就是将所有模块的标志递归地更改为 true。见代码:github.com/pytorch/pytorch/blob/…

标签: python pytorch


【解决方案1】:

model.train() 告诉您的模型您正在训练模型。如此有效的层,如 dropout、batchnorm 等,它们在火车上表现不同,测试程序知道发生了什么,因此可以相应地表现。

更多详情: 它将模式设置为训练 (见source code)。您可以致电model.eval()model.train(mode=False) 来告知您正在测试。 期望 train 函数来训练模型有点直观,但它并没有这样做。它只是设置模式。

【讨论】:

  • 是否有一个标志来检测模型是否处于评估模式?例如mdl.is_eval()?
  • 使用model.training 标志。在eval 模式下是False
  • 在当前的文档中,我发现这个“model.train()”不再被使用:pytorch.org/tutorials/beginner/basics/quickstart_tutorial.html我做了一个小的 3 层神经网络模型的小测试,带有批量规范和 dropout并在表格数据集上对其进行训练。我发现添加 model.train() 实际上阻止了我的模型准确率超过 70%。当我去掉这条线时,准确率是 87%!
  • @Indrajit 您是否检查过它不在训练模型中,即model.trainingFalse?我认为默认情况下这是真的,这就是为什么他们省略了model.train() 电话。至于您的结果,如果不知道数据是什么以及您是否测量测试或训练准确性等,我不能说太多。
  • @UmangGupta - 默认情况下 model.training 是 True,但如果你查看链接,他们的训练循环,在训练步骤之后他们有一个 eval 步骤 - 他们称之为模型。评估()。这将使 model.training 为 False 他们不会重置。我知道这很违反直觉——我也很困惑。仍在试图了解为什么会发生这种情况。
【解决方案2】:

有两种方法可以让模型知道您的意图,即您想训练模型还是使用模型进行评估。 在model.train() 的情况下,模型知道它必须学习层,当我们使用model.eval() 时,它表明模型不需要学习任何新内容,并且该模型用于测试。 model.eval() 也是必要的,因为在 pytorch 中,如果我们使用的是 batchnorm,而在测试期间,如果我们只想传递单个图像,如果未指定 model.eval(),pytorch 会抛出错误。

【讨论】:

    【解决方案3】:

    这里是module.train()的代码:

    def train(self, mode=True):
            r"""Sets the module in training mode."""      
            self.training = mode
            for module in self.children():
                module.train(mode)
            return self
    

    这里是module.eval

    def eval(self):
            r"""Sets the module in evaluation mode."""
            return self.train(False)
    

    模式traineval 是我们可以设置模块的仅有的两种模式,它们完全相反。

    这只是一个 self.training 标志,目前只有 DropoutBatchNorm 关心该标志。

    默认情况下,此标志设置为True

    【讨论】:

    • 现在还有其他支持self.training标志的层吗?
    • 我想知道model.eval() 是如何影响向后传球的?
    • model.eval() 只是一个不采用 dropout 和 batch 规范的开关。我有一个很好的intro to PyTorch training,您可以在其中检查前向和后向传递,以及deep intro to PyTorch AD,您可以在其中自信地了解 PyTorch AD 的详细信息。
    【解决方案4】:

    当前的official documentation 声明如下:

    这仅对某些模块有任何 [原文如此] 效果。如果它们受到影响,请参阅特定模块的文档以了解其在训练/评估模式下的行为的详细信息,例如Dropout、BatchNorm 等。

    【讨论】:

      【解决方案5】:
      model.train() model.eval()
      Sets your model in training mode i.e.

      BatchNorm layers use per-batch statistics
      Dropout layers activated etc


      Sets your model in evaluation (inference) mode i.e.

      BatchNorm layers use running statistics
      Dropout layers de-activated etc.

      Equivalent to model.train(False).

      注意:这些函数调用都不会向前/向后传递。它们告诉模型如何在运行时采取行动。

      这很重要,因为some modules (layers)(例如DropoutBatchNorm)在训练和推理过程中的行为不同,因此如果在错误的模式下运行,模型会产生意想不到的结果。

      【讨论】:

        【解决方案6】:

        考虑以下模型

        import torch
        import torch.nn.functional as F
        from torch_geometric.nn import GCNConv
        
        class GraphNet(torch.nn.Module):
            def __init__(self, num_node_features, num_classes):
                super(GraphNet, self).__init__()
                self.conv1 = GCNConv(num_node_features, 16)
                self.conv2 = GCNConv(16, num_classes)
        
            def forward(self, data):
                x, edge_index = data.x, data.edge_index
                x = self.conv1(x, edge_index)
                x = F.dropout(x, training=self.training) #Look here
                x = self.conv2(x, edge_index)
                return F.log_softmax(x, dim=1)
        

        在这里,dropout 的功能在不同的操作模式下有所不同。如您所见,它仅在self.training==True 时有效。所以,当你输入model.train() 时,模型的前向函数会执行dropout,否则不会(比如model.eval()model.train(mode=False) 时)。

        【讨论】:

          猜你喜欢
          • 2020-04-03
          • 2020-10-14
          • 2018-08-01
          • 2020-01-23
          • 2019-05-30
          • 2021-07-31
          • 2019-12-05
          • 2014-08-03
          • 2012-02-12
          相关资源
          最近更新 更多