본문 바로가기
C.W.K.
Stream
Lesson 03 of 05 · published

NumPyro로 베이지안 Inference

~8 min · scientific, jax, tutorial

Level 0호기심
0 XP0/73 lessons0/17 achievements
0/100 XP to next level100 XP to go0% complete

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가 가장 사용하기 편해.

Code

import numpyro
import numpyro.distributions as dist
from numpyro.infer import MCMC, NUTS, Predictive
import jax
import jax.numpy as jnp

# Define a Bayesian linear regression model
def linear_regression(x, y=None):
    # Priors
    alpha = numpyro.sample('alpha', dist.Normal(0, 10))
    beta = numpyro.sample('beta', dist.Normal(0, 10))
    sigma = numpyro.sample('sigma', dist.HalfNormal(5))

    # Likelihood
    mu = alpha + beta * x
    numpyro.sample('obs', dist.Normal(mu, sigma), obs=y)

# Generate synthetic data
key = jax.random.key(0)
true_alpha, true_beta = 2.0, 3.5
x = jnp.linspace(-5, 5, 100)
y = true_alpha + true_beta * x + 0.5 * jax.random.normal(key, (100,))

# Run MCMC with NUTS (No U-Turn Sampler)
kernel = NUTS(linear_regression)
mcmc = MCMC(kernel, num_warmup=500, num_samples=1000)
mcmc.run(jax.random.key(1), x=x, y=y)

# Get posterior samples
samples = mcmc.get_samples()
print(f"alpha: {samples['alpha'].mean():.2f} ± {samples['alpha'].std():.2f}")
print(f"beta:  {samples['beta'].mean():.2f} ± {samples['beta'].std():.2f}")
# alpha: 2.00 ± 0.05  (true: 2.0)
# beta:  3.50 ± 0.01  (true: 3.5)

# Make predictions
predictive = Predictive(linear_regression, samples)
predictions = predictive(jax.random.key(2), x=jnp.array([0.0, 1.0, 2.0]))
print(f"Predictions: {predictions['obs'].mean(axis=0)}")
# ≈ [2.0, 5.5, 9.0]
import jax
import jax.numpy as jnp

# Monte Carlo estimation of pi using vmap
def estimate_pi(key, num_samples):
    keys = jax.random.split(key, 2)
    x = jax.random.uniform(keys[0], (num_samples,))
    y = jax.random.uniform(keys[1], (num_samples,))
    inside_circle = (x**2 + y**2) <= 1.0
    return 4.0 * jnp.mean(inside_circle)

# Run 100 independent estimates in parallel
keys = jax.random.split(jax.random.key(0), 100)
estimates = jax.vmap(estimate_pi, in_axes=(0, None))(keys, 10000)
print(f"π ≈ {estimates.mean():.4f} ± {estimates.std():.4f}")
# π ≈ 3.1415 ± 0.0162

External links

Exercise

NumPyro의 NUTS sampler로 단순한 Gaussian 평균 모델을 적합하고 posterior trace를 검사해. NUTS 내부의 그래디언트 평가가 자동으로 JIT 컴파일된다는 점이 어떤 가치를 주는지 설명해.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.