PyTorch FSDP API — 파라미터 샤딩과 CPU 오프로드
PyTorch FSDP API — 파라미터 샤딩과 CPU 오프로드
torch.distributed.fsdp 패키지는 PyTorch에서 Fully Sharded Data Parallel을 구현한 API예요. FullyShardedDataParallel로 모델을 감싸 파라미터·그래디언트·옵티마이저 상태를 데이터 병렬 워커 전반에 샤딩하고, 선택적으로 CPU로 오프로드할 수 있어요.
핵심 API
FullyShardedDataParallel(module, ...): 모델을 감싸 샤딩을 적용해요.auto_wrap_policy:default_auto_wrap_policy처럼 특정 조건(예: 파라미터 수 1억 초과)으로 하위 모듈을 재귀적으로 감싸요. 수동으로는enable_wrap()/wrap()로 원하는 부분만 감쌀 수 있어요.CPUOffload(offload_params=True): 계산에 쓰이지 않는 동안 쪼갠 파라미터를 CPU로 내려 GPU 메모리를 더 아껴요(전송 오버헤드 트레이드오프).
왜 계층으로 감싸나요
레이어를 중첩으로 감싸면 각 FSDP 인스턴스는 자기 레이어 계산에 필요한 파라미터만 한 번에 모으고, 계산이 끝나면 바로 메모리를 놓아줘요. 이렇게 하면 피크 GPU 메모리가 크게 줄어 더 큰 모델·더 큰 배치로 확장할 수 있어요.