Hoppa till innehållet
AI-grafen
G· Frontier Labmodelltraning-finjustering· ca 120 min· volatil — kontrolleras ofta· verifierad 2026-09-20

FSDP, ZeRO och DeepSpeed

Kunna förklara sharding av parametrar, gradienter och optimerartillstånd.

Förkunskaper

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:

StegDelasMinne per GPU (N kort)
DDPinget14 + 14 + 56
ZeRO-1optimerartillstånd14 + 14 + 56/N
ZeRO-2+ gradienter14 + (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:

  1. all-gather lagrets parametrar från alla rangs.
  2. Beräkna framåt.
  3. Släpp parametrarna (bara den egna sharden behålls).
  4. 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

Alla källor och licenser