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

JAX 백엔드

~8 min · backend

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

큰 계산을 한 덩어리로 최적화한다

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 학습 단계를 작성할 때는 핵심 관용구가 돼.

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 행렬곱을 JIT 적용 전후로 실행해 시간을 비교해.

Progress

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

댓글 0

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

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