DDP — 데이터 병렬 학습의 축
DDP — 데이터 병렬 학습의 축
DistributedDataParallel(DDP) 는 PyTorch가 제공하는 가장 널리 쓰이는 분산 학습 모듈이에요. 같은 모델의 복제본을 여러 프로세스에 두고, 데이터를 나눠 처리하면서 그래디언트만 공유하는 구조예요.
동작 원리
각 프로세스는 로컬 모델 복제본을 갖고, 모든 복제본이 동일한 초기 상태에서 시작하도록 rank 0의 state_dict() 를 브로드캐스트해요. 순전파 후 역전파에서 그래디언트가 준비되면, Reducer 가 버킷 단위로 그래디언트를 모아 all_reduce 로 평균을 내요.
알아두면 좋은 점
find_unused_parameters=True면 사용하지 않은 파라미터의 그래디언트도 처리해요.- 버킷팅으로 통신을 최대한 일찍 시작해 속도를 확보해요.
- TorchDynamo와 함께 쓸 땐 DDP 래퍼를 컴파일보다 먼저 적용해요.
시작 예제
프로세스 그룹을 만들고 DDP로 모델을 감싼 뒤 forward/backward/optimizer 스텝을 돌리면 돼요. 로컬 모델만 프로세스별로 다루면 나머지는 DDP가 처리해요.