그라디언트 체크포인팅
그라디언트 체크포인팅 (Gradient checkpointing)
forward pass는 일반적으로 backward pass가 재사용할 수 있도록 모든 중간 활성화(activation)를 캐싱해요. 하지만 활성화는 배치 크기와 시퀀스 길이에 비례해 커져요. 그라디언트 체크포인팅은 특정 활성화만 저장하고 나머지는 버려요. 그래서 backward pass가 필요할 때 버려진 활성화 중 일부를 그때그때 다시 계산해야 해요.
출처: 문서
본문
Normal training:
Forward: [L1]→[L2]→[L3]→[L4] (save ALL activations)
Backward: ←uses cached activations everywhere
Gradient checkpointing:
Forward: [L1]→[L2]→[L3]→[L4] (save only at checkpoints, discard the rest)
Backward: ←reaches L2, recomputes L2→L3 from scratch, uses it, discards it
backward pass가 버려진 활성화를 다시 계산하기 때문에 학습은 보통 더 느려지지만, 체크포인팅은 활성화 메모리를 줄여줘요.
gradient_checkpointing=True로 설정하면 활성화돼요.
[!TIP] gradient accumulation과 함께 사용하면 메모리 사용량을 더 줄일 수 있어요.
from transformers import TrainingArguments
args = TrainingArguments(
...,
gradient_checkpointing=True,
)
부분 체크포인팅 (Partial checkpointing)
전체 그라디언트 체크포인팅은 체크포인팅 가능한 모든 레이어를 다시 계산해요. 실행에 여유 메모리가 조금 있다면, 메모리 절약 일부를 속도와 맞바꿔서 더 적은 레이어를 체크포인팅할 수 있어요.
every_n_layers를 gradient_checkpointing_enable()에 전달해서 체크포인팅 간격을 선택해요.
every_n_layers=2
Forward: input -> [L1] -> [L2] -> [L3] -> [L4] -> [L5] -> [L6]
CP keep CP keep CP keep
Backward: output <- [L6] <- [L5] <- [L4] <- [L3] <- [L2] <- [L1]
keep rerun keep rerun keep rerun
every_n_layers=2라면 첫 번째 레이어와 그 뒤로 매 두 번째 레이어가 체크포인팅돼요. 체크포인팅된 레이어는 forward pass에서 자신의 활성화를 버리고 backward pass에서 다시 계산하는 반면, 나머지 레이어는 활성화를 메모리에 유지해요.
model.gradient_checkpointing_enable(every_n_layers=2)
기본값인 every_n_layers=1은 모든 레이어를 체크포인팅해요. 더 큰 값은 첫 번째 레이어와 그 뒤로 매 n번째 레이어를 체크포인팅하고, 나머지 레이어의 활성화는 메모리에 남겨둬요. 예를 들어 every_n_layers=2는 레이어 1, 3, 5 등을 체크포인팅해요. GradientCheckpointingLayer를 상속받는 모듈만 계산에 포함돼요. 그라디언트 체크포인팅을 지원하는 다른 모듈은 계속 활성화된 상태로 유지돼요.
Trainer로 부분 그라디언트 체크포인팅을 사용하려면 gradient_checkpointing_kwargs에 every_n_layers를 설정해요.
from transformers import TrainingArguments
args = TrainingArguments(
...,
gradient_checkpointing=True,
gradient_checkpointing_kwargs={"every_n_layers": 4},
)
저장된 활성화 offloading
그라디언트 체크포인팅은 체크포인팅된 각 레이어당 하나의 활성화를 GPU에 유지해요. 시퀀스 길이가 길어지면 상당한 메모리를 소모할 수 있어요 (대략 layers x sequence x hidden x bytes_per_element). offload를 설정하면 그 활성화들을 pinned host 메모리에 보관해요. 대신 forward pass에서 device-to-host 복사가, backward pass에서 host-to-device 복사가 일어나요.
from transformers import TrainingArguments
args = TrainingArguments(
...,
gradient_checkpointing=True,
gradient_checkpointing_kwargs={"offload": True},
)
두 복사 모두 compute stream에서 실행되므로, 메모리를 얻는 대가로 더 느린 step이 되죠. 실행이 다른 방법으로는 안 들어가지 않을 때 사용하고, 실행을 빠르게 하려고 쓰는 건 아니에요.
더 알아보기 (Learn more)
- 학습 중 GPU 메모리를 무엇이 사용하고 있는지 이해하려면 GPU memory usage 문서를 읽어 보세요.
- 더 낮은 정밀도 데이터 타입으로 메모리를 더 줄이고 학습을 빠르게 하는 방법은 Mixed precision training 가이드를 확인해 보세요.
- 커스텀 fused kernel로 학습을 빠르게 하는 방법은 Kernels 가이드를 확인해 보세요.