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

JAX 학습 루프가 명시적인 이유

~8 min · training, jax, tutorial

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

PyTorch Lightning, Keras의 model.fit(...) 같은 상위 수준 래퍼가 JAX 코어에는 없어. 모든 학습 스텝을 직접 작성해. 처음에는 답답하지만 의도된 설계야.

JAX 식 학습 루프의 모양

@jax.jit
def train_step(state, batch):
    x, y = batch

    def loss_fn(params):
        pred = model.apply(params, x)
        loss = compute_loss(pred, y)
        metrics = {"acc": accuracy(pred, y)}
        return loss, metrics

    (loss, metrics), grads = jax.value_and_grad(loss_fn, has_aux=True)(state.params)
    updates, new_opt_state = optimizer.update(grads, state.opt_state, state.params)
    new_params = optax.apply_updates(state.params, updates)

    new_state = state.replace(
        params=new_params,
        opt_state=new_opt_state,
        step=state.step + 1,
    )
    return new_state, loss, metrics

# 사용자가 직접 loop
for batch in dataloader:
    state, loss, metrics = train_step(state, batch)
    if state.step % 100 == 0:
        print(f"step {state.step}: loss={loss:.4f}")

30줄이면 충분하고 모든 과정이 보여. 마법처럼 숨은 곳도 없어.

왜 이게 좋은가?

  • 모든 step이 가시적: 그래디언트, 옵티마이저 상태, 매개변수 갱신을 모두 직접 확인할 수 있어. 학습기.fit() 안에 숨어 있던 부분이 모두 드러나.
  • 커스터마이징 자유: 그래디언트 자르기, custom 갱신 규칙, 다양한 스케줄, 모두 같은 30줄 안에 추가해. 래퍼 API의 한계 없어.
  • 디버깅 단순: print, assert, breakpoint, 어디든 자유. 마법 같은 callback 시스템 없어.
  • JIT 명확: train_step 함수가 정확히 무엇을 jit하는지 드러나고 컴파일 비용도 확인할 수 있어.

비교, PyTorch Lightning

# PyTorch Lightning
class MyModel(pl.LightningModule):
    def training_step(self, batch, batch_idx):
        x, y = batch
        pred = self.model(x)
        loss = F.cross_entropy(pred, y)
        self.log("train_loss", loss)
        return loss

    def configure_optimizers(self):
        return torch.optim.AdamW(self.parameters(), lr=1e-3)

trainer = pl.Trainer(max_epochs=10, gpus=4)
trainer.fit(model, dataloader)

편리하지만 무엇이 어떻게 돌아가는지 알려면 Lightning의 소스 코드를 살펴봐야 해. JAX에서는 직접 쓴 30줄이 곧 소스 코드야.

고수준 래퍼가 필요할 때

Flax의 nnx.training, chex, clu 같은 라이브러리가 일부 상용구를 줄여 줘. 그러나 핵심은 같은 명시적 패턴이고 보조 도구만 추가해. PyTorch Lightning 같은 통합 래퍼가 JAX 코어에 없는 건 의도된 설계야.

🛠 JAX 학습 코드의 mantra

"Show me the 루프." JAX 코드를 받으면 train_step 함수 한 개 + 루프 한 개. 30줄로 끝이야. 래퍼가 늘어날수록, 그 코드가 JAX 답지 않는지 의심해 봐. 학습 코드가 짧아지는 것보다, 명료한 게 더 가치 있다는 게 JAX 공동체의 가치관.

한 가지, 명시적 학습 코드는 처음 익히는 데 학습 곡선이 있어. 몇 번 짜 보면 패턴이 눈에 들어와서, 새 작업마다 빠르게 응용할 수 있어. PyTorch Lightning의 콜백들을 외우는 것보다, 다른 작업에도 더 잘 옮겨 갈 수 있는 지식이야.

Code

import jax
import jax.numpy as jnp

# The pattern: forward → loss → grad → update → repeat
def train_step(params, x, y, lr=0.01):
    # 1. Define loss as a function of params
    def loss_fn(params):
        predictions = model_forward(params, x)
        return jnp.mean((predictions - y) ** 2)

    # 2. Compute loss and gradients simultaneously
    loss, grads = jax.value_and_grad(loss_fn)(params)

    # 3. Update parameters (simple SGD)
    new_params = jax.tree.map(lambda p, g: p - lr * g, params, grads)

    return new_params, loss

# 4. JIT-compile for speed
train_step_jit = jax.jit(train_step)

# 5. Training loop
for epoch in range(num_epochs):
    for batch in data_loader:
        params, loss = train_step_jit(params, *batch)

External links

Exercise

JAX에서 학습 반복문을 직접 작성하는 방식과 PyTorch Lightning의 configure_optimizers, training_step 방식을 비교해. 장단점을 다섯 개 항목으로 정리하고, JAX가 포기하게 하는 것과 돌려주는 것을 한 문장으로 요약해.

Progress

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

댓글 0

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

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