긴 시퀀스나 깊은 모델 학습의 한 가지 큰 적, 역전파에 필요한 순전파 활성값이 메모리를 대부분 차지해. 해결: 그래디언트 체크포인팅. 계산을 더 하는 대신 메모리를 절약해.
표준 backprop의 메모리 패턴
forward: x → h1 → h2 → h3 → h4 → loss
[저장] [저장] [저장] [저장]
backward: 모든 h 사용해서 grad 계산
L계층 모델은 모든 계층의 활성값을 메모리에 보존해. 시퀀스 길이가 4096인 Transformer 계층 24개라면 메모리 사용량이 폭발해.
체크포인팅은 활성값의 일부 또는 전부를 저장하지 않고 역전파에서 다시 계산해
import jax
def expensive_layer(params, x):
# 큰 activation 만드는 layer
h = jnp.tanh(x @ params["W1"])
h = h @ params["W2"]
return h
# checkpoint 적용 — forward 에선 activation 안 저장, backward 에서 재계산
checkpointed = jax.checkpoint(expensive_layer)
def model(params, x):
for layer_params in params:
x = checkpointed(layer_params, x)
return x
# 학습 — 메모리 절감, compute 추가
loss, grads = jax.value_and_grad(loss_fn)(params, x, y)
Transformer 학습에서는 보통 메모리를 50% 절약해 두 배 큰 배치를 사용할 수 있는 대신, 순전파를 한 번 더 수행해 연산량이 약 33% 늘어.
세분성: 어디까지 체크포인팅할까?
# 전체 model — 너무 거침
checkpointed_model = jax.checkpoint(model)
# 각 layer — 표준
def model(params, x):
for layer in params:
x = jax.checkpoint(layer_fn)(layer, x)
return x
# 매 N layer 마다 — 더 미세 조정
N = 4
def model(params, x):
for i in range(0, len(params), N):
chunk = params[i:i+N]
x = jax.checkpoint(lambda c, x: chunk_fn(c, x))(chunk, x)
return x
가장 좋은 체크포인트 단위는 모델마다 달라. attention 같은 큰 계층은 따로 체크포인팅하고 작은 연산은 묶어.
policy를 직접 지정
import jax.checkpoint_policies as ckpt_policies
# 중요한 op (matmul 같은 거) 만 저장, 나머지는 재계산
checkpointed = jax.checkpoint(
expensive_layer,
policy=ckpt_policies.checkpoint_dots_with_no_batch_dims,
)
# 또는 직접
checkpointed = jax.checkpoint(
expensive_layer,
policy=ckpt_policies.dots_saveable,
)
실전: Transformer 학습
def transformer_block(params, x, mask):
'''attention + MLP — 큰 activation'''
h = layer_norm(x, params["ln1"])
h = attention(params["attn"], h, mask)
x = x + h
h = layer_norm(x, params["ln2"])
h = mlp(params["mlp"], h)
x = x + h
return x
# 각 block 마다 checkpoint
def model(params, x, mask):
for block_params in params["blocks"]:
x = jax.checkpoint(transformer_block)(block_params, x, mask)
x = layer_norm(x, params["final_ln"])
return x @ params["head"]
같은 GPU에서 네 배 긴 시퀀스를 학습할 수 있지만 학습 속도는 약 30% 느려지는 절충이 있어.
⚖️ 메모리 vs 연산
학습 메모리의 대부분은 순전파 활성값이 차지해. jax.checkpoint는 그 메모리 비용을 계산 시간으로 바꿔. 큰 모델이나 긴 시퀀스일수록 효과가 크고, 모델에서 OOM이 나면 가장 먼저 시도할 만한 최적화야. ZeRO 같은 sharding보다 진입 장벽이 낮고 효과도 즉각적이야.
JAX의 jax.remat는 jax.checkpoint의 별칭이야. remat은 옛 이름이고, 최근 코드는 checkpoint가 표준이야.