C.W.K.
Stream
Lesson 05 of 08 · published

JAX Backend

~8 min · backend

Level 0Keras 도제
0 XP0/97 lessons0/20 achievements
0/120 XP to next level120 XP to go0% complete

속도·스케일의 backend

JAX backend (JAX 0.4.20+) 는 GPU·TPU 학습에서 자주 *가장 빠른 선택*. op 하나하나 dispatch 하는 대신 XLA 가 계산 전체를 하나의 fused·하드웨어 튜닝된 kernel 로 컴파일하거든. Google large-scale 학습 (Gemini, PaLM 등) 의 엔진이고, model 이 진지해지면 의미 생기는 셋을 Keras 로 들고 와:

  • XLA JIT 컴파일로 throughput 극대화
  • 함수형 변환 — jax.vmap·jax.grad·jax.pmap — 네 math 위에 조합 가능
  • keras.distribution (multi-GPU/TPU data·model 병렬) 지원 최강

functional purity 라는 세금

JAX 의 힘은 제약에서 나와: 변환되는 함수는 *pure* 해야 해 — 숨은 mutable state 금지. 근데 Keras model 객체는 state (weight, optimizer slot, metric 집계) 를 들고 있잖아. 그래서 JIT/grad 변환 안에서 돌리려면 in-place 로 바꾸는 대신 state 를 넣고 *업데이트된 state 를 돌려받아*. 그게 Code block 의 stateless_call 패턴이야. fit() 만 쓰면 평생 안 만질 수도 있는데, custom JAX train step 짜는 순간 핵심 관용구가 돼.

Code

JAX stateless call — state 넣고 돌려받기·python
os.environ["KERAS_BACKEND"] = "jax"
import keras

model = keras.Sequential([...])
model.compile(optimizer="adam", loss="mse")

# JAX stateless API for functional purity
variables = model.variables
outputs = model.stateless_call(variables, inputs)

External links

Exercise

KERAS_BACKEND=jax 로 작은 keras.ops.matmul 스니펫을 @jax.jit 데코레이터 함수 안에 넣어. 1000×1000 matmul 의 jit 유/무 시간 비교.

Progress

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

댓글 0

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

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