속도·스케일의 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 짜는 순간 핵심 관용구가 돼.