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

JAX가 과학 계산에서 빛나는 이유

~8 min · scientific, jax, tutorial

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

JAX는 머신러닝 프레임워크로도 이미 강력해. 하지만 진짜 차별점은 머신러닝 밖의 영역에서 드러나. JAX가 차별화되는 이유는 수치 코드를 자유롭게 미분하고 가속할 수 있기 때문이야.

전통적인 과학 계산 도구:

  • NumPy / SciPy, 주로 CPU에서 실행되고 일반 수치 코드의 자동 미분은 지원하지 않아.
  • MATLAB, 도구 상자가 풍부하지만 폐쇄형이고 그래디언트는 따로 작성해야 해.
  • Fortran / C++, 빠르지만 그래디언트를 직접 작성해야 해.
  • Numba는 JIT을 지원하지만 자동 미분은 지원하지 않아.

JAX가 합치는 두 가지:

  1. 임의 수치 코드의 자동 미분: ODE 풀이, 시뮬레이션, 샘플링, 수치 코드를 미분할 수 있어.
  2. 가속기 활용: 같은 Python 코드를 GPU와 TPU에서 실행할 수 있어.

이 둘이 합쳐지면 새 카테고리의 알고리즘이 가능해져.

전형적 사용 사례

1. 미분 가능한 물리 시뮬레이션: 시뮬레이터를 통해 그래디언트를 흘려 물리 시스템을 학습해.

def simulate_pendulum(initial_angle, length, dt, n_steps):
    '''단순한 진자 시뮬레이션'''
    theta, omega = initial_angle, 0.0
    for _ in range(n_steps):
        omega += -9.81 / length * jnp.sin(theta) * dt
        theta += omega * dt
    return theta

# 진자가 특정 각도에 도달하도록 길이를 학습
target_angle = 0.5
def loss(length):
    final = simulate_pendulum(0.1, length, 0.01, 100)
    return (final - target_angle) ** 2

# 자동 미분 — 시뮬레이션 통해 gradient 흐름
optimal_length = optimize_with_grad(loss)

2. 미분방정식: Diffrax는 미분 가능한 ODE와 SDE solver를 제공해.

from diffrax import diffeqsolve, Tsit5, ODETerm

def lorenz(t, y, args):
    x, y_, z = y
    return jnp.array([
        10 * (y_ - x),
        x * (28 - z) - y_,
        x * y_ - 8/3 * z,
    ])

solution = diffeqsolve(
    ODETerm(lorenz),
    Tsit5(),
    t0=0.0, t1=10.0, dt0=0.01,
    y0=jnp.array([1., 1., 1.]),
)

3. 베이지안 추론: NumPyro는 JAX 위에서 확률적 프로그래밍을 제공해.

import numpyro
import numpyro.distributions as dist

def model(data):
    mu = numpyro.sample("mu", dist.Normal(0, 1))
    sigma = numpyro.sample("sigma", dist.HalfNormal(1))
    with numpyro.plate("data", len(data)):
        numpyro.sample("obs", dist.Normal(mu, sigma), obs=data)

# NUTS sampler — JAX-jit 으로 자동 가속
mcmc = numpyro.infer.MCMC(numpyro.infer.NUTS(model), num_samples=1000)
mcmc.run(rng_key, data)

4. 최적화: JAXopt와 Lineax 같은 라이브러리를 사용할 수 있어.

import jaxopt

# 비선형 least squares — 자동 미분 + GPU
solver = jaxopt.LevenbergMarquardt(residual_fun=residuals)
result = solver.run(init_params, data=observations)

5. 확률적 프로그래밍, TensorFlow Probability JAX 백엔드

6. 강화학습, Brax (미분할 수 있는 물리 시뮬레이션), MJX (MuJoCo)

7. 양자 컴퓨팅, qujax, PennyLane JAX 백엔드

8. 계산화학, JAX-MD (molecular dynamics)

🔬 미분할 수 있는 시뮬레이션의 의미

전통적인 과학에서는 관측값으로부터 모델 매개변수를 추정해. JAX에서는 시뮬레이션 자체를 미분할 수 있어. 그래서 그래디언트 하강법을 이용한 매개변수 학습, 실험 설계, 민감도 분석을 하나의 프레임워크에서 다룰 수 있어. "무엇이든 미분할 수 있다"는 관점은 새로운 연구 문제를 여는 강력한 패러다임이야.

이어지는 레슨에서는 Diffrax의 ODE, NumPyro의 베이지안 추론, Brax와 JAX-MD의 시뮬레이션을 통해 JAX 생태계의 각 영역을 살펴봐.

Code

import jax
import jax.numpy as jnp

# Example: differentiable physics simulation
def simulate_spring(k, m, x0, v0, dt, num_steps):
    """Simulate a damped spring system: m*x'' + 0.1*x' + k*x = 0"""
    def step(state, _):
        x, v = state
        a = (-k * x - 0.1 * v) / m  # spring force + damping
        v_new = v + a * dt
        x_new = x + v_new * dt
        return (x_new, v_new), x_new

    init_state = (x0, v0)
    _, trajectory = jax.lax.scan(step, init_state, None, length=num_steps)
    return trajectory

# Simulate
traj = simulate_spring(k=2.0, m=1.0, x0=1.0, v0=0.0, dt=0.01, num_steps=1000)

# Gradient: how does the final position change with spring constant?
@jax.jit
def final_position(k):
    return simulate_spring(k, 1.0, 1.0, 0.0, 0.01, 1000)[-1]

dk = jax.grad(final_position)(2.0)
print(f"d(final_pos)/dk = {dk:.6f}")

# Vectorize: simulate 100 different spring constants at once
ks = jnp.linspace(0.5, 5.0, 100)
all_trajectories = jax.vmap(lambda k: simulate_spring(k, 1.0, 1.0, 0.0, 0.01, 1000))(ks)
print(f"Batch trajectories shape: {all_trajectories.shape}")  # (100, 1000)

External links

Exercise

미분 가능성과 가속이 전통적인 과학 Python 도구인 NumPy, SciPy, Numba보다 유리한 분야 다섯 곳을 찾아 목록으로 만들어. 그중 직접 써 볼 분야 하나를 고르고 세 줄로 이유를 적어. 목록보다 문제를 바라보는 틀이 더 중요해.

Progress

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

댓글 0

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

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