🧮 Computational Mathematics

Inicio · Parte 15 — Matemática de Deep Learning

319 — Autodiff con PyTorch/JAX

deep-learning clase 19 de 20 4 horas demostración autodiff_frameworks

PyTorch y JAX hacen lo mismo que el Var de la parte 08, con ingeniería de por medio.

Fórmulas

modo reverso: una pasada adelante, una atrás
coste ≈ 2 veces el del forward, independientemente del número de parámetros
modo directo: eficiente con pocas entradas y muchas salidas

Desarrollo

La autodiferenciación en modo reverso obtiene los gradientes de una salida escalar respecto de todas las entradas con una sola pasada hacia atrás. Su coste es aproximadamente el doble del paso hacia adelante, independientemente de cuántos parámetros haya. Esa propiedad es lo que hace viable entrenar modelos de miles de millones de parámetros.

El modo directo propaga derivadas hacia adelante y es eficiente en el caso opuesto: pocas entradas y muchas salidas. Como en aprendizaje automático la pérdida siempre es un escalar y los parámetros son millones, el modo reverso es el adecuado, y por eso es el que implementan todos los frameworks.

Ninguno de los dos es diferenciación simbólica ni numérica. No manipula fórmulas ni usa diferencias finitas: evalúa derivadas exactas de operaciones elementales y las compone según el grafo. Es exacta hasta el redondeo y eficiente, que es lo mejor de ambos mundos.

Lo que aportan PyTorch y JAX sobre el Var de la parte 08 no es el concepto sino la ingeniería: núcleos optimizados para GPU y TPU, fusión de operaciones, compilación diferida, paralelismo y una cobertura enorme de operaciones. El principio se entiende en cien líneas de Python; la implementación de producción son cientos de miles.

Ejemplo trabajado

La misma expresión derivada por el motor propio.

expresión: loss = (tanh(wx + b) − 1)²

loss = 0,09543807

dloss/dw = −0,48417723
dloss/db = −0,32278482
dloss/dx = −0,22594937

Tres gradientes de una sola pasada hacia atrás.

Con un millón de parámetros el coste sería el mismo
factor 2 sobre el forward: esa es la propiedad clave.

Con diferencias finitas harían falta un millón de
evaluaciones adicionales, una por parámetro.

Qué calcula el laboratorio

Nuestro Var frente a PyTorch/JAX: mismo principio, distinta escala.

python classes/part-15-matematica-de-deep-learning/319-autodiff-con-pytorch-jax/lab.py
compmath run 319

Salidas del laboratorio (9)

Muestra de la ejecución real

{
  "expresion": "loss = (tanh(wx + b) - 1)²",
  "loss": 0.09543807,
  "dloss/dw": -0.48417723,
  "dloss/db": -0.32278482,
  "dloss/dx": -0.22594937,
  "frameworks_disponibles": {
    "torch": false,
    "jax": false,
    "numpy": false
  }
}

Errores comunes

Dónde se usa

Todo entrenamiento moderno, optimización de simuladores diferenciables, física diferenciable y cálculo de sensibilidades.

Idea rectora de la parte

Normalizar estabiliza la escala interna y permite tasas de aprendizaje mayores.

Error a evitar

Inicializar todos los pesos iguales y romper la simetría nunca.

Conexión con IA

Toda arquitectura moderna, incluido el Transformer, se construye sobre estos bloques y sobre este mismo mecanismo de derivación.

Bibliografía de la clase

Archivos de la clase