El problema que me hizo buscar una alternativa.
. Mi trabajo implica tomar modelos del Universo (ecuaciones de estado de energía oscura, gravedad modificada, campos taquiónicos) y preguntar: ¿qué dicen realmente los datos sobre los parámetros? La herramienta para esa pregunta es la inferencia bayesiana. Por lo general, realizo un muestreo anidado de dinastia para evaluaciones de probabilidad de unos pocos miles a unos cientos de miles, dependiendo de la complejidad del modelo.
Durante la mayor parte de mi doctorado, no pensé mucho en el solucionador de ODE dentro de la probabilidad de que solve_ivp funcionara. Fue confiable. Por eso lo usé y seguí adelante.
Luego comencé a trabajar en un modelo taquiónico de energía oscura DBI en el que el campo de energía oscura está gobernado por un término cinético no estándar, y las ecuaciones de fondo y de perturbación son un sistema rígido acoplado. Cada llamada de probabilidad resolvió esas EDO, calculó la distancia de comovimiento y evaluó el módulo de distancia en los desplazamientos al rojo de 30 supernovas.
Lo perfilé. La resolución ODE por sí sola tardaba 0,4 ms por llamada. En una ejecución de muestreo anidada con 10⁵ evaluaciones, eso son 40 segundos, solo en llamadas ODE, antes de contar cualquier contabilidad. Y para un modelo de 10 parámetros, obtener un gradiente a través de diferencias finitas centrales cuesta 20 soluciones directas adicionales, convirtiendo esos 0,4 ms en 8 ms por gradiente. Eso son 300 segundos, o unos 5 minutos, sólo para los gradientes. Para una única ejecución de muestreo anidada.
Algo tenía que cambiar.
Lo que encontré: difrax
Después de un día de búsqueda, llegué a Difrax.[1], una biblioteca de solucionadores de ODE numéricos escritos íntegramente en JAX. No es un sustituto neuronal. No es una aproximación. Los mismos algoritmos integrados de Runge-Kutta que ya uso en scipy (Tsit5 en lugar de RK45, pero la misma familia de métodos), simplemente compilados, diferenciables y vectorizables.
Tres propiedades provienen del diseño “escrito íntegramente en JAX”:
Compilación JIT: todo el ciclo de pasos adaptativos se compila en un único kernel XLA. Cero sobrecarga de Python después de la primera llamada.
Autodiff: debido a que cada operación dentro del solucionador es una primitiva JAX, jax.grad propaga gradientes a través de la resolución. Degradados exactos. Un pase hacia atrás. Independientemente de cuántos parámetros.
vmap: se puede resolver un lote completo de vectores de parámetros en paralelo con jax.vmap. Crítico para el muestreo anidado.
Instalarlo lleva 10 segundos:
pip instalar jax difrax
El problema de la prueba: ΛCDM plano de supernovas
Para concretar la comparación, permítanme mostrarles el problema exacto con el que estaba trabajando. En un universo ΛCDM plano, la distancia de comovidad satisface:
dχdz=cH(z),H(z)=H0Ωm(1+z)3+(1−Ωm),χ(0)=0frac{dchi}{dz} = frac{c}{H(z)}, quad H(z) = H_0sqrt{Omega_m(1+z)^3 + (1-Omega_m)}, quad chi(0)=0
El módulo de distancia es el siguiente: μ(z) = 5 log₁₀[(1+z)χ(z) / 10 pc]. Quiero inferir (Ωₘ, H₀) a partir de 30 observaciones simuladas del módulo de distancia SNIa.
de scipy.integrate import solve_ivp import numpy as np C_KMS = 299792.458 # velocidad de la luz [km/s] def rhs(z, chi, Om, H0): return C_KMS / (H0 * np.sqrt(Om*(1+z)**3 + (1-Om))) def forward_scipy(Om, H0, z_obs): sol = solve_ivp(rhs, t_span=(0, z_obs[-1]), y0=[0.0], t_eval=z_obs, args=(Om, H0), método="RK45", rtol=1e-8, atol=1e-10) chi = sol.y[0]devolver 5 * np.log10((1 + z_obs) * chi * 1e5) # módulo de distancia
A la antigua usanza: SciPy
de scipy.integrate import solve_ivp import numpy as np C_KMS = 299792.458 # velocidad de la luz [km/s] def rhs(z, chi, Om, H0): return C_KMS / (H0 * np.sqrt(Om*(1+z)**3 + (1-Om))) def forward_scipy(Om, H0, z_obs): sol = solve_ivp(rhs, t_span=(0, z_obs[-1]), y0=[0.0], t_eval=z_obs, args=(Om, H0), método="RK45", rtol=1e-8, atol=1e-10) chi = sol.y[0]devolver 5 * np.log10((1 + z_obs) * chi * 1e5) # módulo de distancia
La nueva forma: Diffrax
import jax, jax.numpy as jnp import difrax as dfx # No negociable: habilitar 64 bits (más sobre esto a continuación) jax.config.update("jax_enable_x64", True) def H_jax(z, Om, H0): return H0 * jnp.sqrt(Om*(1+z)**3 + (1-Om)) @jax.jit # compilar una vez, llamar rápido para siempre def forward_diffrax(theta, z_obs): Om, H0 = theta[0], theta[1]sol = dfx.diffeqsolve( dfx.ODETerm(lambda z, chi, a: C_KMS / H_jax(z, a[0], a[1])), dfx.Tsit5(), t0=0.0, t1=float(z_obs[-1]), # valor inicial y final dt0=1e-3, # tamaño de paso inicial y0=jnp.array(0.0), # condición inicial args=(Om, H0), saveat=dfx.SaveAt(ts=z_obs), stepsize_controller=dfx.PIDController(rtol=1e-8, atol=1e-10), max_steps=10_000, ) chi = sol.ys return 5 * jnp.log10((1 + z_obs) * chi * 1e5)
La física es idéntica. El algoritmo de resolución es casi idéntico (Tsit5 es muy similar a RK45). Las únicas diferencias estructurales son @jax.jit y la API difrax. Veamos qué aportan esos dos cambios.
Sorpresa 1: la velocidad
solve_ivp: 404 μs por llamada. Difrax post-JIT: 59 μs por llamada. Eso es 07 veces más rápido.
Me quedé mirando este número durante unos segundos la primera vez que lo vi. Permítanme ser honesto acerca de dónde viene realmente la aceleración, porque no es mágica.
En solve_ivp, Python vuelve a ingresar al backend de C/Cython en cada llamada. La memoria se asigna nueva. El bucle while adaptativo pasa por el intérprete de Python y pregunta: "¿el error local es demasiado grande? Rechace; de lo contrario, aumente el paso; repita". Para una solución de 12 pasos, es decir, 12 rondas de envío de Python, 12 asignaciones, 12 cálculos de estimación de errores detrás del bloqueo del intérprete.
En difrax, la primera llamada @jax.jit rastrea todo el cálculo, incluido el bucle while adaptativo, que se reduce a lax. while_loop y se lo entrega a XLA para que lo compile en un núcleo de código de máquina. Cada llamada posterior ejecuta ese kernel directamente. Por lo tanto, no hay Python, no hay necesidad de asignación ni envío.
Para 100.000 evaluaciones de probabilidad, 404 μs frente a 59 μs se traducen en 40,4 segundos frente a 5,9 segundos. Esa es la diferencia que se mejora cuando aumenta la complejidad del modelo.
Sorpresa 2: los gradientes se vuelven gratuitos
Esta fue la parte que cambió no sólo mi flujo de trabajo sino también mi forma de pensar sobre la inferencia. Con scipy, obtener un gradiente de probabilidad logarítmica con respecto a 2 parámetros (Ωₘ, H₀) cuesta 4 soluciones directas (diferencias finitas centrales). Una vez que empiezas a subir el dial, se vuelve caro rápidamente: 10 parámetros significan 20 soluciones directas, 50 parámetros significan 100. La factura crece linealmente con el número de parámetros.
∂ℱ∂Ωm≈ℱ(Ωm+h,H0)−ℱ(Ωm−h,H0)2h,∂ℱ∂H0≈ℱ(Ωm,H0+h)−ℱ(Ωm,H0−h)2hfrac{partialmathcal{F}}{partialOmega_m} approx frac{mathcal{F}(Omega_m+h,H_0) – mathcal{F}(Omega_m-h,H_0)}{2h}, qquad frac{partialmathcal{F}}{partial H_0} approx frac{mathcal{F}(Omega_m,H_0+h) – mathcal{F}(Omega_m,H_0-h)}{2h}
Con difrax escribo:
def loss(theta): mu_pred = forward_diffrax(theta, z_obs) return 0.5 * jnp.sum(((mu_pred – mu_obs) / sigma_mu)**2) grad_fn = jax.jit(jax.grad(loss)) # ese es el cambio completo g = grad_fn(jnp.array([0.3, 70.0])) # gradiente exacto
Debajo del capó, el autodiff de modo inverso de JAX integra las ecuaciones adjuntas[2]hacia atrás a través de la resolución ODE, pero nunca tengo que escribir esas ecuaciones. El resultado es un gradiente exacto en el tiempo comparable a un paso hacia adelante, independientemente del número de parámetros.
Cómo elegir un solucionador
Cuando se trata de elegir un solucionador, hay que tener un poco de cuidado. Opté por Tsit5 por defecto para casi todo y resolvió aproximadamente el 95% de mis problemas sin quejarme. Si quieres todo el proceso de decisión, aquí lo tienes:
EDO no rígida (la mayoría de los problemas cosmológicos) → dfx.Tsit5() ← comience aquí Tolerancias muy ajustadas (< 10⁻⁹) → dfx.Dopri8() EDO rígida (muchos pasos, el solucionador parece lento) → dfx.Kvaerno5() Términos rígidos + no rígidos (IMEX) → dfx.KenCarp4() SDE → dfx.EulerHeun() o dfx.SPaRK()
Una forma rápida de saber si su problema es difícil: imprima sol.stats["num_steps"]. Si es entre 10 y 100 veces mayor de lo esperado, el problema es complicado y necesita un solucionador implícito.
La recompensa: la inferencia cosmológica de un extremo a otro
Ahora, permítanme mostrarles la comparación de inferencias completa. Empiezo ambas canalizaciones desde la misma suposición inicial errónea (Ωₘ, H₀) = (0,10, 60), muy alejada de la verdad (0,30, 70), y ejecuto 350 pasos de gradiente.
Tubería scipy: gradiente a partir de diferencias finitas centrales, descenso de gradiente simple, tasa de aprendizaje fija. Canalización de difrax: gradiente de autodiff, optimizador Adam con un programa de tasa de aprendizaje de desintegración del coseno. importar optax # optimizadores para JAX # Escalar parámetros para que Adam pueda manejarlos por igual # Om ~ 0.3, h = H0/100 ~ 0.7 — ambos O(1) ahora def loss_scaled(theta_s): theta = jnp.array([theta_s[0], 100,0 * theta_s[1]]) return loss(theta) grad_scaled = jax.jit(jax.grad(loss_scaled)) Schedule = optax.cosine_decay_schedule( init_value=0.05, decay_steps=350, alpha=0.04) opt = optax.adam(schedule) theta = jnp.array([0.10, 0.60]) # empezar lejos de la verdad state = opt.init(theta) para paso en rango(350): g = grad_scaled(theta) actualizaciones, state = opt.update(g, state) theta = optax.apply_updates(theta, actualizaciones) if (paso + 1) % 50 == 0: print(f"Paso {paso+1}: Om={theta[0]:.3f} H0={100*theta[1]:.2f}")
Si bien el oleoducto difrax recupera parámetros físicamente sensibles, el oleoducto scipy no puede mover ambos parámetros simultáneamente: un error de manual en el descenso de gradiente en problemas mal escalados. Adam maneja esto automáticamente a través de sus tasas de aprendizaje adaptativo por parámetro, pero Adam solo está disponible porque autodiff me brinda gradientes exactos para alimentarlo.
Tres cosas en las que me equivoqué (para que tú no tengas que hacerlo)
Advertencia 1: olvidar la precisión de 64 bits. JAX tiene como valor predeterminado flotantes de 32 bits. Si se superan las tolerancias (rtol < 10⁻⁷), se pueden obtener resultados muy extraños: en mi ODE, el solucionador necesita 69 pasos en 32 bits, pero sólo 12 en 64 bits. Si ajusta aún más las tolerancias, puede fallar por completo. La solución es simple: habilite 64 bits antes de hacer cualquier otra cosa:
jax.config.update("jax_enable_x64", True) # debe ser el primero
Advertencia 2: evaluación comparativa sin calentamiento. La primera llamada a cualquier función decorada con @jax.jit incluye un hit de compilación único de aproximadamente 90 a 100 ms. Si incluye eso en sus tiempos, difrax parecerá más lento que scipy por el motivo equivocado. La solución es calentar una vez y desechar esa primera carrera:
_ = forward_diffrax(theta, z_obs).block_until_ready() # compilar # AHORA punto de referencia: esta es la velocidad real
Además: JAX se envía de forma asincrónica. Siempre llame a .block_until_ready() en ciclos de tiempo o mida el tiempo para enviar el trabajo, no terminarlo.
Advertencia 3: la trampa del orden de los argumentos. scipy.odeint espera f(y, t) (primero el estado, segundo el tiempo). Casi todo lo demás (solve_ivp, difrax) espera f(t, y). Si portas el antiguo código odeint a difrax sin intercambiar los argumentos, terminarás resolviendo una ODE diferente y normalmente no obtendrás un error. Simplemente obtendrás la respuesta incorrecta.
¿Deberías hacer el cambio?
La respuesta honesta es esta: si estás resolviendo una ODE única y no necesitas gradientes, solve_ivp está perfectamente bien; no hay necesidad de aprender una nueva API. Pero si está haciendo inferencias (evaluaciones de probabilidad repetidas, gradientes de parámetros o soluciones por lotes), vale la pena el esfuerzo.
La migración en sí es pequeña. El modelo avanzado cambia en unas seis líneas. El degradado aparece añadiendo una línea más. El resto del código de inferencia permanece idéntico.
Una cosa que debemos mencionar aquí es que difrax no está "basado en ML" en el sentido de utilizar una red neuronal. Son las mismas matemáticas clásicas de Runge-Kutta, escritas en JAX. La "aceleración ML" proviene de la compilación JIT y la autodiff, ambas herramientas de infraestructura del mundo ML aplicadas a un solucionador numérico clásico. El único enfoque genuinamente basado en ML sería un sustituto neuronal que aprenda θ → μ(z) a partir de datos de entrenamiento, un tema aparte y más avanzado.
El código de trabajo completo.
Todo lo anterior en un script independiente (pip install jax diffrax optax):
""" flat_lcdm_inference.py Infiere (Omega_m, H0) de 30 supernovas simuladas usando difrax + Adam. pip install jax diffrax optax """ import jax, jax.numpy as jnp, numpy as np import diffrax as dfx, optax from scipy.integrate import solve_ivp # solo para generar datos simulados jax.config.update("jax_enable_x64", True) # — Constantes y datos ———————————————– C_KMS = 299792.458 z_obs = jnp.linspace(0.05, 1.5, 30) SIGMA = 0.10 # Datos simulados en verdad (Om=0.30, H0=70) def chi_np(Om, H0): sol = solve_ivp(lambda z, y: C_KMS/(H0*np.sqrt(Om*(1+z)**3+(1-Om))), (0, 1.5), [0.], t_eval=np.array(z_obs), rtol=1e-10) return sol.y[0]mu_true = 5*np.log10((1+np.array(z_obs))*chi_np(0.3, 70.)*1e5) mu_obs = jnp.array(mu_true + 0.10*np.random.default_rng(42).standard_normal(30)) # — modelo directo de difrax ——————————————– @jax.jit def adelante(theta): Om, H0 = theta[0], theta[1]sol = dfx.diffeqsolve( dfx.ODETerm(lambda z, chi, a: C_KMS/(a[1]*jnp.sqrt(a[0]*(1+z)**3+(1-a[0])))), dfx.Tsit5(), t0=0., t1=1.5, dt0=1e-3, y0=jnp.array(0.), args=(Om, H0), saveat=dfx.SaveAt(ts=z_obs), stepsize_controller=dfx.PIDController(rtol=1e-8, atol=1e-10), max_steps=10_000, ).ys return 5*jnp.log10((1+z_obs)*sol*1e5) # — Pérdida y gradiente ———————————————— def loss(th_s): # optimizar en coordenadas escaladas (Om, h=H0/100) mu = forward(jnp.array([th_s[0], 100.*th_s[1]])) return 0.5*jnp.sum(((mu – mu_obs)/SIGMA)**2) grad_fn = jax.jit(jax.grad(loss)) # Calentar el compilador JIT theta_init = jnp.array([0.10, 0.60]) _ = forward(jnp.array([0.3, 0.7])).block_until_ready() _ = grad_fn(theta_init).block_until_ready() # — Optimizador Adam con programación coseno LR ————————— sched = optax.cosine_decay_schedule(init_value=0.05, decay_steps=350, alpha=0.04) opt = optax.adam(sched) theta = theta_init state = opt.init(theta) print(f"{'Step':>5} {'Om':>7} {'H0':>7} {'Loss':>8}") para el paso en el rango (350): g = grad_fn(theta) upd, state = opt.update(g, state) theta = optax.apply_updates(theta, upd) if (step + 1) % 70 == 0 o paso == 0: L = float(pérdida(theta)) print(f"{paso+1:5d} {float(theta)[0]):7.4f} {100*flotante(theta[1]):7.3f} {L:8.2f}") Om_fit, H0_fit = flotante(theta[0]), 100*flotante(theta[1]) print(f"nFinal: Om = {Om_fit:.3f} H0 = {H0_fit:.2f}") print(f"Verdad: Om = 0.300 H0 = 70.00")
Números de un vistazo
El resultado de scipy "incorrecto" no es una falla del solucionador; refleja que el simple descenso de gradiente con gradientes de diferencia finita no puede manejar el desajuste de escala de 200 × entre Ωₘ y H₀.
Pensamiento final
Cambiar mi modelo directo a difrax no cambió la física ni el método de inferencia. Cambió la viabilidad práctica de hacer esa inferencia. Una ejecución de muestreo anidado que se dirigía hacia un presupuesto de modelo futuro de gran tiempo se convirtió en una ejecución de menos de minutos. Los gradientes que iban a costar 20 soluciones adicionales por paso se volvieron esencialmente gratuitos.
La curva de aprendizaje fue de aproximadamente una tarde. La depuración fue principalmente la advertencia de 64 bits y la confusión del calentamiento JIT. La recompensa ha sido real e inmediata.
Si es físico y utiliza scipy para evaluaciones de probabilidad repetidas y aún no ha analizado difrax, espero que esto le dé una razón para hacerlo.
Una nota sobre la reproducibilidad: los tiempos exactos que vea diferirán en su máquina e incluso entre ejecuciones en la misma máquina. En mi Mac (modelo base MacBook Air M3), la llamada de avance de difrax varió entre 55 µs y 62 µs entre sesiones, y scipy varió entre 400 µs y 407 µs. Esto es normal: el estado térmico de la CPU, la programación del sistema operativo y las condiciones de la memoria caché cambian los números absolutos entre un 10% y un 15%. Lo que se mantiene estable es la proporción: difrax es consistentemente entre 07 y 08 veces más rápido que scipy en este problema. La proporción, no el tiempo absoluto, es el número que hay que sacar.
El código Python que generó cada figura de este artículo está disponible en: github.com/Samit1424/ODE_solver_comparison
Nota: Excluyendo la imagen destacada, que se produjo con una herramienta de inteligencia artificial, todas las ilustraciones son obra original del autor.
Referencias
[1]P. Kidger, On Neural Differential Equations, tesis de doctorado, Universidad de Oxford, 2021. docs.kidger.site/diffrax/
[2]RTQ Chen, Y. Rubanova, J. Bettencourt, D. Duvenaud, Ecuaciones diferenciales neuronales ordinarias, NeurIPS 2018.