DataParallel보다 DDP가 유리한 이유
nn.DataParallel은 한 프로세스 안에서 batch를 GPU별로 나누고 결과를 0번 GPU에 모은다.
1
| model = torch.nn.DataParallel(model)
|
사용은 간단하지만 큰 모델과 데이터에서는 0번 GPU에 작업이 몰리고 GPU 사이 불균형이 생길 수 있다. DistributedDataParallel(DDP)은 GPU마다 프로세스를 두고 동일한 모델로 각자 batch를 처리한다.
원문 전체 예제는 pytorch_multi_gpu 저장소에 있다. 아래는 구조를 읽기 위한 핵심 조각이며 dataset, model, optimizer, train 함수가 빠져 있어 완전한 실행 코드가 아니다.
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
| def main():
world_size = torch.cuda.device_count()
torch.multiprocessing.spawn(
main_worker,
nprocs=world_size,
args=(world_size,),
)
def main_worker(gpu, world_size):
torch.distributed.init_process_group(
backend='nccl',
init_method='tcp://127.0.0.1:3456',
world_size=world_size,
rank=gpu,
)
torch.cuda.set_device(gpu)
model = model.cuda(gpu)
model = torch.nn.parallel.DistributedDataParallel(
model, device_ids=[gpu]
)
sampler = torch.utils.data.distributed.DistributedSampler(
train_dataset
)
loader = torch.utils.data.DataLoader(
train_dataset,
shuffle=False,
sampler=sampler,
)
|
DistributedSampler를 쓸 때 shuffle=False로 두는 이유는 DataLoader와 sampler가 동시에 순서를 섞지 않게 하기 위해서다. global batch를 GPU 수로 나눌지, GPU당 batch를 유지해 global batch를 키울지도 명확히 정해야 한다.
checkpoint는 0번 프로세스에서만 저장한다.
모든 프로세스가 동시에 같은 파일을 쓰면 손상되거나 서로 덮어쓸 수 있다. 반대로 load는 각 프로세스가 자신의 GPU 위치로 map해 같은 상태를 가져와야 한다. dist.barrier()는 모든 프로세스가 그 지점에 들어올 때까지 기다리므로, 위치를 잘못 잡으면 느림이나 멈춤처럼 보일 수 있다.