Curso / PyTorch / Guardar, cargar y compilar
● piloto de formato

PyTorch · Fundamentos

Guardar, cargar, y el cambio de reglas de PyTorch 2.0

La parte aburrida (persistir un modelo) y la parte que cambió todo (compilarlo) — ambas resuelven el mismo problema de fondo: sacarle a Python de encima lo que no necesita estar ahí.

1. state_dict vs. guardar el modelo completo

Un modelo entrenado guarda sus parámetros aprendidos en un diccionario interno, el state_dict — mapea el nombre de cada parámetro a su tensor de valores. Es la forma recomendada de persistir un modelo:

torch.save(model.state_dict(), 'model_weights.pth')

# para cargar: primero instanciar la MISMA clase, después cargar los pesos
model = models.vgg16()
model.load_state_dict(torch.load('model_weights.pth', weights_only=True))
model.eval()   # no lo olvides — ver la página anterior

Notá que cargar el state_dict requiere tener antes una instancia de exactamente la misma clase de modelo — el state_dict son solo números, no sabe nada de la arquitectura que los usa. La alternativa, torch.save(model, 'model.pth'), guarda el objeto completo (arquitectura incluida) usando pickle de Python — más cómodo, pero acopla el archivo guardado a que la definición exacta de la clase siga existiendo y siendo importable cuando se cargue de vuelta. La documentación oficial es clara: state_dict es la práctica recomendada; guardar el objeto completo es "un caso de uso legado".

2. El problema que PyTorch cargó durante 5 años: eager es lento

La fortaleza histórica de PyTorch —ejecución eager, cada línea corre inmediatamente en Python, sin un grafo estático de por medio— es también su límite de performance. Desde que se lanzó PyTorch en 2017, el hardware (GPUs) se volvió ~15× más rápido en cómputo, pero mantener esa velocidad en modo eager obligó a mover cada vez más lógica interna de Python a C++, lo cual va en contra del valor central de PyTorch: que el código sea hackeable y fácil de extender.

El equipo de PyTorch probó varios enfoques de compilación a lo largo de 5 años — torch.jit.trace, TorchScript, FX tracing, Lazy Tensors— y ninguno resolvía el problema completo: algunos eran flexibles pero lentos, otros rápidos pero exigían reescribir el modelo de forma no trivial.

3. torch.compile: un decorador, no una reescritura

torch.compile (PyTorch 2.0, 2023) resolvió esto sin pedirle al usuario que cambie una sola línea de su modelo:

model = torch.compile(model)   # eso es todo — el resto del código no cambia

Por dentro corren cuatro tecnologías nuevas, cada una resolviendo una parte distinta del problema de compilación:

  • TorchDynamo: captura el grafo de cómputo de forma segura usando un hook del intérprete de CPython (PEP-0523) — sobre 7.000+ proyectos reales de GitHub usados como validación, acertó a capturar el grafo correctamente el 99% de las veces, contra menos del 50% de TorchScript.
  • AOTAutograd: genera de antemano ("ahead-of-time") el grafo del backward pass, no solo el forward.
  • PrimTorch: reduce los más de 2000 operadores de PyTorch a un conjunto cerrado de ~250 operaciones primitivas — simplifica muchísimo escribir un backend nuevo.
  • TorchInductor: el compilador final, que genera código Triton para GPU y C++/OpenMP para CPU.

4. Los números reales, y el matiz importante

Sobre un benchmark de 163 modelos open-source (46 de Hugging Face Transformers, 61 de TIMM, 56 de TorchBench), sin modificar el código de esos modelos más allá de envolverlos en torch.compile:

MétricaResultado medido
Tasa de éxito de compilación93% de los modelos
Speedup en FP32, A100+21% en promedio
Speedup con mixed precision (AMP), A100+51% en promedio

Sylvain Gugger, mantenedor principal de Hugging Face Transformers en ese momento, lo resumió así: "con solo una línea de código, PyTorch 2.0 da un speedup de entre 1.5x y 2x entrenando modelos Transformer — lo más emocionante desde que se introdujo mixed precision training". El matiz honesto que la propia documentación marca: los speedups son menores en GPUs de consumo (una RTX 3090, por ejemplo) que en GPUs de servidor como A100, y TorchInductor no soporta todavía todas las arquitecturas de GPU por igual.

🔧 El mismo espíritu que ya conocés de NVFP4

torch.compile y la cuantización NVFP4 de tu pipeline atacan el mismo problema desde ángulos distintos: ambos buscan que el hardware haga más trabajo útil por unidad de tiempo/memoria, sin que el usuario tenga que reescribir el modelo. Uno lo hace generando mejor código (kernels fusionados, menos overhead de Python); el otro, usando menos bits por número. En producción real, ambos se combinan: un modelo cuantizado y compilado no son técnicas competidoras, son ortogonales.

5. Resumen

  1. state_dict es la forma recomendada de guardar un modelo — solo números, requiere reinstanciar la clase para cargar; guardar el objeto completo con pickle es un caso de uso legado.
  2. La ejecución eager de PyTorch es flexible pero deja rendimiento sobre la mesa — el problema que PyTorch 2.0 resolvió después de 5 años de intentos previos.
  3. torch.compile no requiere reescribir el modelo — TorchDynamo captura el grafo con 99% de éxito sin tocar el código original.
  4. Los números reales: 93% de modelos compilan sin problema, con speedups de 21% (FP32) a 51% (mixed precision) en GPUs de servidor.
  5. Compilación y cuantización no compiten — atacan ejes distintos del mismo problema (velocidad vs. memoria) y se combinan en producción real.