DDP — 데이터 병렬 학습의 축

DDP — 데이터 병렬 학습의 축

DistributedDataParallel(DDP) 는 PyTorch가 제공하는 가장 널리 쓰이는 분산 학습 모듈이에요. 같은 모델의 복제본을 여러 프로세스에 두고, 데이터를 나눠 처리하면서 그래디언트만 공유하는 구조예요.

동작 원리

각 프로세스는 로컬 모델 복제본을 갖고, 모든 복제본이 동일한 초기 상태에서 시작하도록 rank 0의 state_dict() 를 브로드캐스트해요. 순전파 후 역전파에서 그래디언트가 준비되면, Reducer 가 버킷 단위로 그래디언트를 모아 all_reduce 로 평균을 내요.

알아두면 좋은 점

  • find_unused_parameters=True 면 사용하지 않은 파라미터의 그래디언트도 처리해요.
  • 버킷팅으로 통신을 최대한 일찍 시작해 속도를 확보해요.
  • TorchDynamo와 함께 쓸 땐 DDP 래퍼를 컴파일보다 먼저 적용해요.

시작 예제

프로세스 그룹을 만들고 DDP로 모델을 감싼 뒤 forward/backward/optimizer 스텝을 돌리면 돼요. 로컬 모델만 프로세스별로 다루면 나머지는 DDP가 처리해요.

더 알아보기