¿Cuándo debo usar .eval() ? Entiendo que se supone que me permite "evaluar mi modelo". ¿Cómo lo vuelvo a apagar para entrenar?
Ejemplo de código de entrenamiento usando .eval() .
model.eval() es una especie de interruptor para algunas capas/partes específicas del modelo que se comportan de manera diferente durante el tiempo de entrenamiento e inferencia (evaluación). Por ejemplo, Dropouts Layers, BatchNorm Layers, etc. Debe desactivarlas durante la evaluación del modelo y .eval() lo hará por usted. Además, la práctica común para evaluar/validar es usar torch.no_grad() junto con model.eval() para desactivar el cálculo de gradientes:
# evaluate model: model.eval() with torch.no_grad(): ... out_data = model(data) ... PERO, no olvide volver al modo de training después del paso de evaluación:
# training step ... model.train() ...model.train() | model.eval() |
|---|---|
| Establece el modelo en modo de entrenamiento : • capas de normalización 1 usan estadísticas por lote • activa las capas de Dropout 2 | Establece el modelo en modo de evaluación (inferencia): • las capas de normalización usan estadísticas en ejecución • desactiva las capas de Dropout |
Equivalente a model.train(False) . |
Puede desactivar el modo de evaluación ejecutando model.train() . Debe usarlo cuando ejecute su modelo como un motor de inferencia, es decir, cuando pruebe, valide y prediga (aunque prácticamente no hará ninguna diferencia si su modelo no incluye ninguna de las capas que se comportan de manera diferente ).
BatchNorm , InstanceNormmodel.eval es un método de torch.nn.Module :
eval()Establece el módulo en modo de evaluación.
Esto tiene algún efecto solo en ciertos módulos. Consulte la documentación de módulos particulares para obtener detalles de sus comportamientos en el modo de capacitación/evaluación, si se ven afectados, por ejemplo,
Dropout,BatchNorm, etc.Esto es equivalente a
self.train(False).
El método opuesto es model.train explicado muy bien por Umang Gupta.