프로덕션 학습 코드에는 빠른 스텝 반복을 위한 scan과 안정적인 체크포인트를 위한 Orbax가 필요해.
jax.lax.scan으로 학습 스텝 반복
Python for 루프는 매 스텝마다 호스트와 장치가 통신하고 Python 오버헤드가 생겨. scan은 모든 스텝을 하나의 jit IR로 묶어 가속기에서 한 번에 실행해.
@jax.jit
def train_n_steps(state, batches, n):
'''n 개 step 을 scan 으로 한 번에'''
def body(state, batch):
new_state, loss, metrics = train_step(state, batch)
return new_state, (loss, metrics)
final_state, (losses, metrics_arr) = jax.lax.scan(body, state, batches)
return final_state, losses, metrics_arr
# 사용
N_STEPS = 1000
all_batches = make_batches(N_STEPS) # shape: (N_STEPS, B, ...)
state, losses, metrics = train_n_steps(state, all_batches, N_STEPS)
print(f"평균 loss: {losses.mean():.4f}")
1000스텝을 한 번의 jit 호출로 처리하므로 호스트 오버헤드가 거의 없어. 다만 모든 배치가 미리 장치에 있어야 하므로 큰 데이터셋은 chunk로 나눠야 해.
현실적인 패턴은 chunk 단위로 scan을 실행하고 chunk 사이는 Python으로 잇는 거야:
CHUNK = 100 # 100 step 단위로 scan
for epoch in range(num_epochs):
for chunk_idx in range(50): # 50 chunk = 5000 step
chunk_batches = get_next_chunk(CHUNK)
state, losses, metrics = train_n_steps(state, chunk_batches, CHUNK)
print(f"chunk {chunk_idx}: avg loss = {losses.mean():.4f}")
Orbax 체크포인트
큰 모델 학습은 몇 시간에서 며칠이 걸려. 중간에 중단됐다고 처음부터 다시 시작하면 안 되므로 Orbax 체크포인트를 사용하는 게 표준이야.
pip install orbax-checkpoint
import orbax.checkpoint as ocp
from etils import epath
# checkpoint manager 설정
ckpt_dir = epath.Path("/tmp/my_model_ckpts")
options = ocp.CheckpointManagerOptions(
save_interval_steps=500, # 500 step 마다 저장
max_to_keep=3, # 최근 3 개만 보관
)
mgr = ocp.CheckpointManager(
ckpt_dir,
item_names=("state",),
options=options,
)
# 저장
mgr.save(
step=state.step,
args=ocp.args.Composite(state=ocp.args.StandardSave(state)),
)
mgr.wait_until_finished()
# 복원 (latest)
restored = mgr.restore(
mgr.latest_step(),
args=ocp.args.Composite(state=ocp.args.StandardRestore(state)),
)
state = restored["state"]
print(f"복원: step {state.step}")
완전한 학습 루프 + 체크포인트
def train_with_resume(initial_state, resume=True):
state = initial_state
if resume and mgr.latest_step() is not None:
restored = mgr.restore(
mgr.latest_step(),
args=ocp.args.Composite(state=ocp.args.StandardRestore(state)),
)
state = restored["state"]
print(f"resumed from step {state.step}")
for epoch in range(num_epochs):
for chunk_idx in range(num_chunks):
chunk_batches = get_next_chunk(CHUNK)
state, losses, metrics = train_n_steps(state, chunk_batches, CHUNK)
# checkpoint 자동 저장
mgr.save(
step=state.step,
args=ocp.args.Composite(state=ocp.args.StandardSave(state)),
)
mgr.wait_until_finished()
return state
프로세스가 중단된 뒤 다시 시작하면 스텝 카운터, 옵티마이저 상태, 매개변수를 마지막 체크포인트부터 자동으로 복원해.
멀티호스트 체크포인트
큰 모델을 멀티호스트로 학습할 때는 모든 호스트가 같은 체크포인트를 봐야 해. Orbax가 이를 자동으로 처리해:
# 모든 host 가 같은 mgr instance 만들고, save 호출
# Orbax 가 host 0 만 disk write, 나머지는 wait
# distributed file system (GCS, S3) 권장
🎯 학습 안정성 체크리스트
(1) scan + jit, 호스트 오버헤드 제거해. (2) Orbax 체크포인트, 매 N 스텝 자동으로 처리해. (3) max_to_keep, 디스크가 가득 안 차게. (4) wait_until_finished, exit 전 호출해. (5) 최신 체크포인트에서 재개하기, 프로세스 재시작 자동으로 처리해. (6) 체크포인트 안에 난수 키도, 학습 재현성 유지해. (7) 작은 실험에서 복원 시험, 첫 코드 짤 때 한 번 실험.
이 패턴이 Llama / Gemini / GPT 학습 코드의 표준 골격. 모델 크기가 1000배 커도, 같은 7가지 구성에 같은 체크포인트 패턴이야.