¿Llama a forward() en nn.Module ? Pensé que cuando llamamos al modelo, se está utilizando el método forward . ¿Por qué necesitamos especificar train()?
model.train() le dice a su modelo que está entrenando al modelo. De manera efectiva, las capas como abandono, norma de lote, etc., que se comportan de manera diferente en el tren y los procedimientos de prueba, saben lo que está sucediendo y, por lo tanto, pueden comportarse en consecuencia.
Más detalles: Establece el modo de entrenar (ver código fuente ). Puede llamar a model.eval() o model.train(mode=False) para indicar que está probando. Es algo intuitivo esperar que la función de train entrene el modelo, pero no lo hace. Simplemente establece el modo.
Aquí está el código de 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 Y aquí está el module.eval .
def eval(self): r"""Sets the module in evaluation mode.""" return self.train(False) Los modos train y eval son los únicos dos modos en los que podemos configurar el módulo, y son exactamente opuestos.
Eso es solo un indicador de self.training y actualmente solo Dropout yBatchNorm se preocupan por ese indicador.
De forma predeterminada, este indicador se establece en True .
model.train() | model.eval() |
|---|---|
| Establece el modelo en modo de entrenamiento , es decir • Las capas BatchNorm usan estadísticas por lote• Capas Dropout activadas, etc. | Establece el modelo en modo de evaluación (inferencia), es decir • Las capas BatchNorm usan estadísticas en ejecución• Capas de Dropout desactivadas, etc. |
Equivalente a model.train(False) . |
Nota: ninguna de estas llamadas de función se ejecuta hacia adelante o hacia atrás. Le dicen al modelo cómo actuar cuando se ejecuta.
Esto es importante ya que algunos módulos (capas) (por ejemplo, Dropout , BatchNorm ) están diseñados para comportarse de manera diferente durante el entrenamiento frente a la inferencia y, por lo tanto, el modelo producirá resultados inesperados si se ejecuta en el modo incorrecto.