PyTorch DDP — 분산 데이터 병렬 학습의 표준

PyTorch DDP — 분산 데이터 병렬 학습의 표준

DistributedDataParallel(DDP) 은 PyTorch에서 제공하는 분산 학습 핵심 모듈이에요. 모델을 DDP로 감싸면 내부적으로 그래디언트 동기화(all-reduce)를 처리해요. 특히 torchrun과 함께 단일 노드·멀티노드 모두에서 안정적으로 동작해요.

어떻게 동작하나요

  • 프로세스당 하나의 모델 복제본이 각자 다른 데이터 배치를 처리해요.
  • 각 순방향·역방향 후에 그래디언트를 all-reduce하여 모든 복제본이 동일한 가중치를 유지해요.
  • param_groups 동기화 등을 자동화해서 코드를 단순하게 만들어요.

torch.distributed.run/torchrun 명령으로 프로세스 그룹을 띄우고, 초기화 후 각 랭크가 자기 데이터 파티션을 학습해요.

언제 좋을까요

배치 크기를 GPU 수만큼 키워도 되는 큰 모델·데이터에서 널리 쓰여요. 다만 모델 자체가 GPU 하나 메모리를 넘어서면 Tensor/파이프라인 병렬로 넘어가야 해요.

더 알아보기