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-토치와 같은 방식으로 쓸 수 있어요.

더 알아보기