FSDP, ZeRO och DeepSpeed
Kunna förklara sharding av parametrar, gradienter och optimerartillstånd.
Förkunskaper
- FDistribuerad träningkrävs
Intuition
Vid dataparallell träning har varje GPU en full kopia av modellen, gradienterna och optimerartillståndet. För en 7B-modell med Adam i bf16 är det ~84 GB per kort — oavsett hur många kort du har.
Sharding delar upp de tre sakerna över korten i stället:
| Steg | Delas | Minne per GPU (N kort) |
|---|---|---|
| DDP | inget | 14 + 14 + 56 |
| ZeRO-1 | optimerartillstånd | 14 + 14 + 56/N |
| ZeRO-2 | + gradienter | 14 + (14+56)/N |
| ZeRO-3 / FSDP | + parametrar | (14+14+56)/N |
Priset är kommunikation: parametrar måste samlas in (all-gather) före varje lagers beräkning och släppas efteråt. Med snabb länk (NVLink, InfiniBand) är det överkomligt; över långsam Ethernet dominerar det.
Formellt
FSDP (PyTorch) och DeepSpeed ZeRO implementerar samma idé med olika API. Flödet i ZeRO-3 per lager:
- all-gather lagrets parametrar från alla rangs.
- Beräkna framåt.
- Släpp parametrarna (bara den egna sharden behålls).
- Samma sak bakåt, plus reduce-scatter av gradienterna.
Kompletterande tekniker:
- Aktiveringscheckpointing: spara inte alla mellanaktiveringar; räkna om dem i bakåtpasset. Sparar mycket minne för ~30 % längre tid. Nästan alltid rätt val vid minnesbrist.
- CPU-offload: flytta optimerartillstånd (och eventuellt parametrar) till värdminnet. Gör mycket stora modeller möjliga på lite GPU-minne, men är långsamt — PCIe-bandbredden blir flaskhals.
- Mixed precision: bf16 för beräkning, fp32 för master-vikter. Standard.
Beslutsordning: börja med DDP. Slut på minne → aktiveringscheckpointing → ZeRO-2 → ZeRO-3/FSDP → offload. Varje steg kostar hastighet, så ta bara nästa när det föregående inte räcker.
Vanligaste praktiska felet: att spara checkpoints från alla rangs (ska vara rank 0, eller state_dict_type=FULL_STATE_DICT med rätt konfiguration) och att glömma sampler.set_epoch() — då ser varje epok likadan ut.
Kod
import torch, functools
import torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP, ShardingStrategy, MixedPrecision
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import apply_activation_checkpointing
dist.init_process_group("nccl")
rank = dist.get_rank(); torch.cuda.set_device(rank)
mp = MixedPrecision(param_dtype=torch.bfloat16, reduce_dtype=torch.bfloat16, buffer_dtype=torch.bfloat16)
wrap = functools.partial(transformer_auto_wrap_policy, transformer_layer_cls={TransformerBlock})
modell = FSDP(modell,
sharding_strategy=ShardingStrategy.FULL_SHARD, # = ZeRO-3
auto_wrap_policy=wrap, # sharda per block, inte hela modellen
mixed_precision=mp,
device_id=rank)
apply_activation_checkpointing(modell, check_fn=lambda m: isinstance(m, TransformerBlock))
auto_wrap_policy är avgörande: utan den shardas modellen som en enhet, vilket innebär att allt samlas in samtidigt — och minnesvinsten uteblir.
Behärskning innebär
- Förklarar sharding av parametrar, gradienter och optimerartillstånd
- Väljer ZeRO-steg efter minnesbehov
- Känner till offloading och aktiveringscheckpointing
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — ZeRO: Memory Optimizations Toward Training Trillion Parameter Models — arXiv (öppen åtkomst; licens per artikel)
- PyTorch — FSDP (BSD-3) — BSD-3-Clause