PyTorch 분산 학습 개요 — DDP·FSDP·TP·PP
PyTorch 분산 학습 개요
PyTorch의 분산 학습은 torch.distributed 를 중심으로 병렬 모듈, 통신 계층, 실행·디버깅 인프라가 하나로 묶인 생태계예요. 처음 분산 학습을 시작한다면 이 개요에서 자신의 상황에 맞는 기술을 고르는 걸 권장해요.
병렬화 모듈
- DistributedDataParallel (DDP): 모델이 GPU 한 장에 들어가지만 여러 GPU로 확장하고 싶을 때.
- FullyShardedDataParallel (FSDP2): 모델이 GPU 한 장에 안 들어갈 때.
- Tensor Parallel (TP) / Pipeline Parallel (PP): FSDP로도 부족할 때 레이어나 텐서를 나눠 적재.
통신 계층 (C10D)
torch.distributed 의 통신 계층은 all_reduce·all_gather 같은 집단 통신 API와 send·isend 같은 P2P API를 제공해요. DDP와 FSDP는 내부적으로 이 계층을 사용해 그래디언트·버퍼를 동기화해요.
실행 도구
torchrun 으로 여러 프로세스를 띄우고, 여러 노드에서 실행할 땐 노드별 설정을 관리해요. 토픽마다 병렬화 기술을 조합하는 게 일반적이에요.