NumPyro는 Pyro의 아이디어를 JAX 위에서 구현한 확률적 프로그래밍 라이브러리야. PyTorch 기반 Pyro가 동적 실행을 택한다면 NumPyro는 JAX의 함수형 설계와 JIT 가속을 활용해.
pip install numpyro
가장 단순, 1D Gaussian mean estimation
import jax
import numpyro
import numpyro.distributions as dist
from numpyro.infer import MCMC, NUTS
def model(data):
'''y ~ Normal(mu, 1), mu ~ Normal(0, 10)'''
mu = numpyro.sample("mu", dist.Normal(0., 10.))
with numpyro.plate("data", len(data)):
numpyro.sample("obs", dist.Normal(mu, 1.), obs=data)
# 합성 데이터
import numpy as np
true_mu = 3.7
data = jax.random.normal(jax.random.PRNGKey(0), (100,)) + true_mu
# MCMC
nuts = NUTS(model)
mcmc = MCMC(nuts, num_warmup=500, num_samples=1000)
mcmc.run(jax.random.PRNGKey(1), data=data)
# 결과
samples = mcmc.get_samples()
print(f"posterior mu: mean={samples['mu'].mean():.3f}, std={samples['mu'].std():.3f}")
print(f"true mu: {true_mu}")
NUTS(No-U-Turn Sampler)는 HMC를 자동으로 조정하는 알고리즘이야. 매 스텝의 그래디언트가 JAX의 자동 미분으로 자동으로 JIT 컴파일돼. 큰 모델에서는 PyTorch 기반 Pyro보다 5~10배 빠른 경우도 흔해.
linear regression
def linear_model(X, y=None):
n_features = X.shape[1]
w = numpyro.sample("w", dist.Normal(0., 1.).expand([n_features]))
b = numpyro.sample("b", dist.Normal(0., 1.))
sigma = numpyro.sample("sigma", dist.HalfNormal(1.))
mean = jnp.dot(X, w) + b
with numpyro.plate("data", len(X)):
numpyro.sample("y", dist.Normal(mean, sigma), obs=y)
# 학습
mcmc = MCMC(NUTS(linear_model), num_warmup=500, num_samples=1000)
mcmc.run(jax.random.PRNGKey(0), X=X_train, y=y_train)
# prediction (posterior predictive)
samples = mcmc.get_samples()
y_pred = jnp.dot(X_test, samples["w"].T) + samples["b"] # (n_samples, n_test)
y_pred_mean = y_pred.mean(0)
y_pred_std = y_pred.std(0)
불확실성을 자연스럽게 추정할 수 있다는 점이 베이지안 방법의 가장 큰 장점이야.
variational inference (VI)
MCMC는 정확하지만 큰 데이터에서는 느리고, VI는 빠른 근사를 제공해.
from numpyro.infer import SVI, Trace_ELBO
from numpyro.infer.autoguide import AutoNormal
import optax
guide = AutoNormal(linear_model)
optimizer = numpyro.optim.optax_to_numpyro(optax.adam(0.01))
svi = SVI(linear_model, guide, optimizer, loss=Trace_ELBO())
svi_result = svi.run(jax.random.PRNGKey(0), 5000, X=X_train, y=y_train)
# guide 로부터 posterior approximation 추출
params = svi_result.params
실전, hierarchical 모델
def hierarchical(group_idx, x, y=None):
'''그룹별 random effect'''
n_groups = len(jnp.unique(group_idx))
# population-level
mu_w = numpyro.sample("mu_w", dist.Normal(0., 1.))
sigma_w = numpyro.sample("sigma_w", dist.HalfNormal(1.))
# group-level
with numpyro.plate("groups", n_groups):
w_g = numpyro.sample("w_g", dist.Normal(mu_w, sigma_w))
# data-level
with numpyro.plate("data", len(x)):
mean = w_g[group_idx] * x
numpyro.sample("y", dist.Normal(mean, 0.1), obs=y)
베이지안 계층 구조를 JAX 위에서 깔끔하게 표현할 수 있고, 큰 데이터와 많은 그룹도 GPU에서 빠르게 처리할 수 있어.
🎲 NumPyro의 가치
베이지안 추론의 큰 비용은 매 샘플링 스텝의 그래디언트 평가와 샘플 사이의 순차 의존성이야. JAX의 jit과 grad가 첫 번째 비용을 자동으로 가속하므로 NUTS의 순차적인 부분만 남아. 덕분에 훨씬 큰 모델과 데이터의 추론도 개인용 노트북에서 시도할 수 있어.
JAX다운 베이지안 도구에는 다른 선택지도 있어. BlackJAX (모듈식 sampler), tfp.substrates.jax (TF Probability의 JAX 백엔드). 각 도구마다 강점이 있지만 NumPyro가 가장 사용하기 편해.