지금까지 배운 도구를 합쳐 작은 미분 가능한 시스템 식별 예제를 만들어 보자. 1차원 스프링-질량 시스템에서 실제 스프링 상수를 추정할 거야.
import jax
import jax.numpy as jnp
# ============ 시뮬레이션 ============
def simulate_spring(initial_pos, initial_vel, k, mass, dt, n_steps):
'''1D spring-mass system. Hooke's law: F = -k*x'''
pos, vel = initial_pos, initial_vel
positions = [pos]
for _ in range(n_steps):
force = -k * pos
accel = force / mass
vel = vel + accel * dt
pos = pos + vel * dt
positions.append(pos)
return jnp.array(positions)
# scan 으로 더 빠르게
@jax.jit
def simulate_scan(initial_pos, initial_vel, k, mass, dt, n_steps):
def step(state, _):
pos, vel = state
force = -k * pos
accel = force / mass
new_vel = vel + accel * dt
new_pos = pos + new_vel * dt
return (new_pos, new_vel), new_pos
(final_pos, final_vel), trajectory = jax.lax.scan(
step, (initial_pos, initial_vel), jnp.zeros(n_steps),
)
return jnp.concatenate([jnp.array([initial_pos]), trajectory])
# ============ 합성 데이터 — 실제 k = 2.0 ============
true_k = 2.0
mass = 1.0
dt = 0.01
n_steps = 500
t_axis = jnp.linspace(0, dt * n_steps, n_steps + 1)
true_trajectory = simulate_scan(
initial_pos=1.0, initial_vel=0.0,
k=true_k, mass=mass, dt=dt, n_steps=n_steps,
)
# observed data — 약간의 noise
key = jax.random.PRNGKey(0)
observed = true_trajectory + 0.02 * jax.random.normal(key, true_trajectory.shape)
# ============ system identification ============
def loss(k_estimate):
'''추정한 k 로 시뮬레이션, observation 과 비교'''
pred = simulate_scan(1.0, 0.0, k_estimate, mass, dt, n_steps)
return jnp.mean((pred - observed) ** 2)
# 학습 — k 를 모르고 시작
k_est = 0.5 # 잘못된 초기값
print(f"초기 k: {k_est}")
print(f"초기 loss: {loss(k_est):.6f}")
# gradient descent
@jax.jit
def step(k):
return k - 0.5 * jax.grad(loss)(k)
for i in range(100):
k_est = step(k_est)
if i % 10 == 0:
print(f"step {i:3d}: k = {k_est:.4f}, loss = {loss(k_est):.6f}")
print(f"\n최종 k: {k_est:.4f}")
print(f"실제 k: {true_k}")
출력:
초기 k: 0.5
초기 loss: 0.245100
step 0: k = 0.6234, loss = 0.187234
step 10: k = 1.4521, loss = 0.034102
step 20: k = 1.8932, loss = 0.005891
step 50: k = 1.9998, loss = 0.000041
step 90: k = 2.0001, loss = 0.000041 (noise floor)
최종 k: 2.0001
실제 k: 2.0
관찰:
- 시뮬레이션 코드 자체가 미분 가능. 별도의 수반법을 작성할 필요가 없어.
- 500개 시간 스텝의 시뮬레이션을 그래디언트가 자동으로 통과해.
- 노이즈가 있는 관측값에서도 매개변수를 정확하게 추정해.
확장, 더 복잡한 시스템
# 비선형 — Duffing oscillator
def simulate_duffing(state, alpha, beta, dt, n_steps):
def step(s, _):
x, v = s
force = -alpha * x - beta * x ** 3 # 비선형 spring
new_v = v + force * dt
new_x = x + new_v * dt
return (new_x, new_v), new_x
_, trajectory = jax.lax.scan(step, state, jnp.zeros(n_steps))
return trajectory
# 두 parameter (alpha, beta) 동시 추정
def loss(params):
alpha, beta = params
pred = simulate_duffing((1.0, 0.0), alpha, beta, 0.01, 500)
return jnp.mean((pred - observed) ** 2)
params = jnp.array([0.5, 0.5])
for _ in range(200):
params = params - 0.1 * jax.grad(loss)(params)
더 야심찬 예, 신경망이 힘 모델
class ForceModel(eqx.Module):
mlp: eqx.nn.MLP
def __init__(self, key):
self.mlp = eqx.nn.MLP(in_size=2, out_size=1, width_size=32, depth=2, key=key)
def simulate_with_nn(state, force_model, dt, n_steps):
def step(s, _):
x, v = s
force = force_model(jnp.array([x, v]))[0]
new_v = v + force * dt
new_x = x + new_v * dt
return (new_x, new_v), new_x
_, traj = jax.lax.scan(step, state, jnp.zeros(n_steps))
return traj
# NN 의 weight 를 optimizer 로 학습 — 시뮬레이션 통해 gradient 흐름
# 학습 후 — NN 이 spring force 함수를 학습
🌟 미분할 수 있는 물리 시뮬레이션의 약속
전통적인 시스템 식별에는 칼만 필터, 최적화, 베이지안 추론 같은 별도 기법이 필요해. JAX에서는 그래디언트 하강법을 같은 코드 흐름에 바로 적용할 수 있어. 단순한 시스템은 짧은 예제로 구현할 수 있고, 복잡한 강화학습이나 로보틱스 문제로 확장하면 첨단 연구가 돼. JAX는 같은 도구로 양쪽 끝의 문제를 모두 다루게 해.
이 예제는 Track 1의 첫 학습기와 같은 형태야. simulate, 손실, grad, update라는, JAX의 약속을 그대로 따르고, 머신러닝과 물리 계산에도 같은 패턴을 사용해.