Empezar por una referencia que responda a la necesidad
Elige un conjunto corto con entradas ordinarias, casos límite y fronteras de decisión. Fija modelo, pesos, preprocesamiento y modo train o eval. Una comparación entre dos modelos o dos batches no permite atribuir su diferencia a la precisión.
Registra la salida útil para tu aplicación, no solo la pérdida. Para un clasificador, puede incluir puntuaciones y decisiones; para una regresión, error y valores extremos. Verifica ya la presencia de NaN o de inf en la referencia. Una ejecución FP32 incorrecta no se convierte en una base fiable porque tenga más bits.
Fija tolerancia, calidad mínima y ausencia de no finitos antes del ensayo. PyTorch recuerda que el cálculo en coma flotante no garantiza resultados idénticos entre dispositivos o rutas de ejecución.
Distinguir autocast, formato numérico y GradScaler
autocast elige el tipo de algunas operaciones según su política de cómputo. No transforma todo el programa a un formato único. Con este uso, evita convertir manualmente todo el modelo con half(). La documentación actual recomienda torch.autocast o torch.amp.autocast; las interfaces antiguas torch.cuda.amp están obsoletas.
GradScaler actúa sobre la escala de la pérdida y de los gradientes durante el entrenamiento. No se emplea como un acelerador de la inferencia, que no realiza backward. FP16 dispone de un rango numérico más restringido que BF16; un modelo diseñado para BF16 puede desbordar en FP16. Una disminución repetida de la escala no establece, por tanto, que el problema esté resuelto.
Elige el formato a partir de las restricciones del modelo y de las operaciones realmente utilizadas, y luego verifica la compatibilidad del destino. El nombre comercial de una tarjeta o una preferencia de preparación de PyTorch no demuestra que tu operador personalizado disponga del kernel deseado.
Desplaza la tabla para leer todas las columnas.| Elección | Rol | Comprobación necesaria |
|---|---|---|
| Referencia FP32 | Punto de comparación del proyecto | Salidas finitas y calidad esperada |
| Autocast FP16 | Algunas operaciones en precisión reducida | Rango numérico y gradientes |
| Autocast BF16 | Otro compromiso rango/precisión | Operadores disponibles y calidad |
| GradScaler | Gestión de la escala de los gradientes | Actualizaciones realmente efectuadas |
Colocar las etapas del entrenamiento en el orden correcto
El fragmento propuesto supone un modelo y un optimizador ya construidos, una entrada y un objetivo en la misma GPU, y una pérdida escalar. No se ha ejecutado y no constituye una validación de una oferta. El contexto autocast rodea el forward y la pérdida; el backward se desarrolla después de su cierre. El scaler se crea una vez para la sesión de entrenamiento, no en cada batch.
Para inspeccionar o recortar los gradientes, elimina primero su factor de escala con unscale_. Los ejemplos oficiales de AMP precisan hacerlo una sola vez por optimizador y después de la acumulación de los gradientes destinados a su actualización. El umbral de recorte 1.0 de abajo es un valor ilustrativo que debes elegir para tu proyecto, no una recomendación universal.
Las guardas interrumpen aquí el diagnóstico si la pérdida, los gradientes o la norma total no son finitos. Estas lecturas en CPU son intrusivas: no cronometres este fragmento. La acumulación, los optimizadores múltiples y el planificador requieren su propia definición de la actualización.
import torch
# Precondiciones: model, optimizer, loss_fn, inputs y targets existen.
# El modelo y las entradas están en el mismo dispositivo CUDA/HIP.
dtype = torch.float16 # Elección a validar; BF16 es otro intento.
scaler = torch.amp.GradScaler("cuda", enabled=(dtype == torch.float16))
# Colocar en tu bucle, con scaler conservado entre los batches.
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type="cuda", dtype=dtype):
prediction = model(inputs)
loss = loss_fn(prediction, targets)
if not bool(torch.isfinite(loss).item()):
raise FloatingPointError("Pérdida no finita: interrumpir el diagnóstico")
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
if any(p.grad is not None and
not bool(torch.isfinite(p.grad).all().item())
for p in model.parameters()):
raise FloatingPointError("Gradiente no finito: interrumpir el diagnóstico")
torch.nn.utils.clip_grad_norm_(
model.parameters(), max_norm=1.0, error_if_nonfinite=True,
)
scaler.step(optimizer)
scaler.update()Ejemplo desarrollado: dos decisiones parecidas no son intercambiables
Supongamos un servicio que elige la clase con la puntuación máxima. Sobre una entrada pedagógica, la referencia produce dos puntuaciones muy próximas: 1,0000 y 1,0003. Otro camino numérico podría modificar su orden o crear un empate. Estos números ilustran una frontera de decisión; no son salidas medidas de FP16 o BF16.
La verificación correcta consta de dos niveles. Compara las puntuaciones con tolerancias explícitas y luego compara la decisión y la regla aplicada a los empates. Una diferencia pequeña en valor absoluto puede cambiar la acción elegida. A la inversa, una diferencia numérica visible puede quedar sin consecuencias para una tarea cuyo umbral se encuentra lejos de las puntuaciones observadas.
Registra identificadores, salidas de referencia, prueba AMP e impacto en la decisión. Fija la regla de aceptación antes de leer los resultados. No amplíes la tolerancia para hacer desaparecer un caso incómodo; salidas de escalas diferentes pueden requerir criterios distintos.
Desplaza la tabla para leer todas las columnas.| Criterio | Referencia | Prueba AMP | Decisión |
|---|---|---|---|
| Salidas finitas | Por verificar | Por verificar | Rechazar los no finitos inexplicados |
| Desviación numérica | Valores conservados | Desviación por calcular | Tolerancia definida antes de la prueba |
| Decisión de aplicación | Clase o acción | Clase o acción | Examinar los cambios |
| Calidad sobre el conjunto fijo | Por medir | Por medir | Respetar el umbral del proyecto |
Interpretar los NaN y las actualizaciones omitidas
Cuando aparezcan no finitos, busca la primera etapa que los produce: entrada, salida intermedia, pérdida o gradiente. Reproduce el mismo caso en referencia y luego desactiva localmente autocast alrededor de la operación sospechosa, controlando también el tipo de sus entradas. Volver a ejecutar todo un entrenamiento en FP32 puede servir de comparación, pero no localiza automáticamente el problema.
El scaler puede evitar una actualización cuando los gradientes contienen inf o NaN. Por lo tanto, un bucle que continúa no necesariamente ha realizado tantas actualizaciones como iteraciones. Registra este comportamiento durante el diagnóstico. No hagas avanzar a ciegas una política de aprendizaje que se supone sigue las actualizaciones efectivas.
Una pérdida finita no garantiza gradientes finitos. A la inversa, un incidente puntual no basta para declarar un entrenamiento inutilizable: examina su frecuencia, la progresión y la calidad. La receta AMP ofrece un método para aislar por separado autocast y el scaling cuando uno de los dos es sospechoso.
Preparar un retorno atrás reproducible
Antes de la prueba, conserva la configuración de referencia, los pesos, el estado del optimizador y un checkpoint coherente. Si tu entrenamiento usa un scaler, su estado también forma parte de la reanudación. Documenta el dtype y las posibles regiones dejadas en FP32. Reanudar con una política diferente es un cambio experimental que debe identificarse, no una continuación implícitamente equivalente.
Vuelve a la configuración anterior si las salidas dejan de ser finitas, si la calidad sale del criterio fijado o si las actualizaciones dejan de avanzar de forma aprovechable. Conserva el caso que motivó ese retorno. Tras una modificación, reinicia la comparación sobre el mismo conjunto antes de extender la duración.
El ejercicio Kernodeck de checkpoints verifica una reanudación en CPU sin AMP. Reutiliza su método de comparación añadiendo los estados que tu bucle consume realmente.
Medir las ganancias solo tras la validación numérica
Tras la validación, mide memoria y tiempo sin el diagnóstico detallado. Mantén formas, batch, modelo y calidad. Las lecturas de escalares, sincronizaciones y perfiladores pueden modificar las duraciones; retira los controles intrusivos de la medición final.
En ROCm, el nombre de dispositivo de PyTorch sigue siendo cuda y las interfaces correspondientes se reutilizan. Esto no garantiza ni los mismos kernels ni resultados idénticos a NVIDIA. Verifica el backend y los operadores del proyecto en el destino. Esta guía no anuncia ninguna reducción fija de memoria, multiplicación de velocidad ni compatibilidad de preparación Kernodeck.