큰 계산을 한 덩어리로 최적화한다
JAX 0.4.20 이상을 백엔드로 쓰면 GPU와 TPU 학습에서 가장 빠른 결과를 얻는 경우가 많아. 연산을 하나씩 실행하는 대신 XLA가 전체 계산을 합치고 하드웨어에 맞춘 커널로 컴파일하기 때문이야. Gemini와 PaLM 같은 Google의 대규모 학습에도 JAX가 쓰였어. Keras에서 JAX를 선택하면 다음 강점을 함께 얻어.
- XLA의 JIT 컴파일로 처리량을 높일 수 있어.
jax.vmap,jax.grad,jax.pmap같은 함수형 변환을 사용자 계산과 조합할 수 있어.keras.distribution을 이용한 다중 GPU·TPU 데이터 및 모델 병렬화 지원이 가장 성숙해.
순수 함수가 요구하는 명시적 상태
JAX의 힘은 변환할 함수가 숨은 가변 상태를 갖지 않아야 한다는 제약에서 나와. 하지만 Keras 모델은 가중치, 옵티마이저 슬롯, 평가지표 누적값 같은 상태를 갖고 있지. JIT나 미분 변환 안에서는 객체를 제자리에서 바꾸는 대신 현재 상태를 인자로 넣고 갱신된 상태를 결과로 받아야 해. 코드 블록의 stateless_call이 그 방식이야. fit()만 쓴다면 직접 만날 일이 드물지만 JAX 학습 단계를 작성할 때는 핵심 관용구가 돼.