DeepSpeed 활성화 체크포인팅 실전
DeepSpeed 활성화 체크포인팅 실전
DeepSpeed는 대규모 분산 훈련에서 torch.utils.checkpoint 의 드롭인 대체 모듈 deepspeed.checkpointing 을 제공해요. 모델 병렬 환경에서 활성화를 더 잘 나누고, CPU 오프로드까지 지원해요.
DeepSpeed만의 확장
- 활성화 분할: 텐서 병렬 랭크 간 활성화를 나눠 GPU당 메모리를 추가로 절감.
- CPU 체크포인팅: 활성화를 GPU HBM 대신 CPU RAM으로 오프로드.
- 연속 메모리 최적화: 활성화를 단일 연속 버퍼에 담아 단편화를 줄임.
- RNG 상태 관리: 재계산 중에 CUDA 랜덤 상태(dropout)를 올바르게 복원.
사용 방법
훈련 시작 전(모델 병렬 MPU 초기화 후) checkpointing.configure(...) 로 설정해요.
import deepspeed.checkpointing as checkpointing
checkpointing.configure(
mpu_=mpu, # 텐서 병렬일 때
partition_activations=True,
contiguous_checkpointing=True,
num_checkpoints=24, # transformer 레이어 수
)
이후 체크포인트하려는 함수를 checkpointing.checkpoint(...) 로 감싸면 pyro-토치와 같은 방식으로 쓸 수 있어요.