Diffrax는 JAX다운 방식으로 설계된 ODE, SDE, CDE solver야. Equinox를 만든 Patrick Kidger가 개발했어. 모든 solver는 jit, grad, vmap과 호환돼.
pip install diffrax
가장 단순한 ODE
import jax.numpy as jnp
from diffrax import diffeqsolve, ODETerm, Tsit5, SaveAt
def vector_field(t, y, args):
'''dy/dt = f(t, y)'''
return -y # exponential decay
solution = diffeqsolve(
ODETerm(vector_field),
solver=Tsit5(), # 5th-order Tsitouras
t0=0.0, t1=5.0, dt0=0.1,
y0=jnp.array(1.0),
saveat=SaveAt(ts=jnp.linspace(0, 5, 100)),
)
print(solution.ts.shape, solution.ys.shape) # (100,) (100,)
Heun, Dopri5, 고차 Dopri8, 강직 문제용 KenCarp 같은 다양한 solver를 같은 API로 사용할 수 있어.
매개변수화된 벡터장
def lorenz(t, y, args):
'''Lorenz attractor'''
sigma, rho, beta = args
x, y_, z = y
return jnp.array([
sigma * (y_ - x),
x * (rho - z) - y_,
x * y_ - beta * z,
])
sol = diffeqsolve(
ODETerm(lorenz),
Tsit5(),
t0=0., t1=10., dt0=0.01,
y0=jnp.array([1., 1., 1.]),
args=(10.0, 28.0, 8/3), # Lorenz parameter
saveat=SaveAt(ts=jnp.linspace(0, 10, 1000)),
)
가장 강력한 점은 그래디언트가 흐른다는 거야
import jax
def simulate(initial_state, params):
sol = diffeqsolve(
ODETerm(lorenz),
Tsit5(),
t0=0., t1=5., dt0=0.01,
y0=initial_state,
args=params,
)
return sol.ys[-1] # 마지막 state
# 시작 state 에 대한 final state 의 gradient
grad_initial = jax.grad(lambda y0: jnp.sum(simulate(y0, params)))
print(grad_initial(jnp.array([1., 1., 1.])))
# ODE solver 를 통해 자동 미분된 gradient
다른 프레임워크에서는 별도의 수반법이나 민감도 코드를 작성해야 할 수 있지만 JAX와 Diffrax에서는 grad를 그대로 적용해. PyTorch torchdiffeq가 비슷하지만, JAX 쪽 사용감이 더 깔끔해.
Neural ODE, 모델이 ODE
class NeuralODE(eqx.Module):
mlp: eqx.nn.MLP
def __init__(self, key):
self.mlp = eqx.nn.MLP(
in_size=3, out_size=3, width_size=64, depth=3, key=key,
)
def __call__(self, t, y, args):
return self.mlp(y)
model = NeuralODE(jax.random.PRNGKey(0))
def integrate(model, y0, t1):
sol = diffeqsolve(
ODETerm(model),
Tsit5(),
t0=0., t1=t1, dt0=0.01,
y0=y0,
)
return sol.ys[-1]
# 학습 — model 의 weight 가 ODE vector field 를 정의
def loss(model, y0, t1, target):
pred = integrate(model, y0, t1)
return jnp.mean((pred - target) ** 2)
grads = jax.grad(loss)(model, y0, 5.0, target)
SDE (확률 미분방정식)
from diffrax import SDESolver, ItoSDE, VirtualBrownianTree
def drift(t, y, args): return -y
def diffusion(t, y, args): return 0.1 * jnp.eye(2)
bm = VirtualBrownianTree(t0=0, t1=10, tol=1e-3, shape=(2,), key=jax.random.PRNGKey(0))
sol = diffeqsolve(
ItoSDE(drift, diffusion),
SDESolver(...),
t0=0., t1=10., dt0=0.01,
y0=jnp.zeros(2),
args=None,
)
🔬 Diffrax의 위치
scipy.integrate는 훌륭한 solver지만 JAX의 자동 미분 시스템과 통합되지는 않아. Diffrax는 JAX다운 방식으로 설계돼 jit, grad, vmap과 호환되고 GPU와 TPU에서 ODE, SDE, CDE를 가속해. Neural ODE, 미분 가능한 물리 시뮬레이션, 민감도 분석의 기반이므로 JAX가 과학 계산에서 빛나는 대표적인 사례야.
같은 개발자가 만든 Optimistix (root finding, fixed-point), Lineax (linear solvers). JAX의 과학 계산 도구 생태계가 빠르게 자라고 있어.