What does model.train() do in PyTorch?

后端 未结 3 885
無奈伤痛
無奈伤痛 2020-12-02 07:15

Does it call forward() in nn.Module? I thought when we call the model, forward method is being used. Why do we need to specify train()

3条回答
  •  悲哀的现实
    2020-12-02 08:13

    model.train() tells your model that you are training the model. So effectively layers like dropout, batchnorm etc. which behave different on the train and test procedures know what is going on and hence can behave accordingly.

    More details: It sets the mode to train (see source code). You can call either model.eval() or model.train(mode=False) to tell that you are testing. It is somewhat intuitive to expect train function to train model but it does not do that. It just sets the mode.

提交回复
热议问题