PyTorch FSDP2 튜토리얼 — 완전 샤딩 데이터 병렬 시작하기
PyTorch FSDP2 튜토리얼 — 완전 샤딩 데이터 병렬 시작하기
Fully Sharded Data Parallel(FSDP) 은 PyTorch가 데이터 병렬 학습을 확장하는 표준 방법이에요. DDP가 각 랭크에 모델 복제본을 통째로 두고 그래디언트를 all-reduce로 합치는 반면, FSDP는 모델 파라미터·그래디언트·옵티마이저 상태 자체를 랭크들에 쪼개서 둬요. 덕분에 한 장비에 안 들어가는 큰 모델도 학습할 수 있어요.
FSDP2 동작 원리
- 순방향·역방향 밖에서는 파라미터가 완전 샤딩돼 있어 메모리를 아껴요.
- 순방향·역방향 직전에 쪼갠 파라미터를 all-gather로 모아 한 데 펼치고(unshard), 계산이 끝나면 다시 쪼개요.
- 역방향에서는 로컬 그래디언트를 reduce-scatter로 쪼갠 그래디언트로 바꿔요.
- 옵티마이저가 쪼갠 그래디언트·상태 그대로 갱신해 메모리를 줄여요.
이 과정을 DDP의 all-reduce 하나를 reduce-scatter + all-gather 둘로 분해한 것으로 이해하면 쉬워요.
코드로 보기
from torch.distributed.fsdp import fully_shard
model = Transformer()
for layer in model.layers:
fully_shard(layer) # 레이어마다 쪼갬
fully_shard(model) # 루트 모델도 쪼갬
optim = torch.optim.Adam(model.parameters(), lr=1e-2)
torchrun --nproc_per_node 2 train.py로 두 프로세스에 걸쳐 실행해요. FSDP1은 deprecated 되고 FSDP2의 fully_shard를 권장해요.