Hoppa till innehållet
AI-grafen
F· AI engineeringmodelltraning-finjustering· ca 90 min· volatil — kontrolleras ofta· verifierad 2026-09-20

Distribuerad träning

Kunna förklara data-, tensor- och pipelineparallellism, och köra träning på flera GPU:er.

Förkunskaper

Intuition

Tre sätt att dela upp träning över flera GPU:er — de löser olika problem:

StrategiDelasLöserKommunikation
Dataparallellismbatchenlångsam träningall-reduce av gradienter varje steg
Tensorparallellismvarje matrismodellen ryms inte på ett kortmycket, inom varje lager → kräver snabb länk
Pipelineparallellismlagrenmodellen ryms intemindre, men ger «bubblor» av väntan
ZeRO / FSDPoptimizer-tillstånd, gradienter, vikterminne vid dataparallellismmer än ren DP, mindre än TP

Regel: börja med dataparallellism (DDP). Räcker inte minnet: lägg på ZeRO/FSDP. Ryms modellen fortfarande inte: tensorparallellism inom en nod, pipeline mellan noder.

Formellt

ZeRO i tre steg (Rajbhandari m.fl. 2019) — varje steg delar upp en sak till över alla N GPU:er:

StegDelar uppMinne per GPU (7B, Adam, bf16)
— (ren DDP)inget14 + 14 + 56 = 84 GB
ZeRO-1optimizer-tillstånd14 + 14 + 56/N
ZeRO-2+ gradienter14 + (14+56)/N
ZeRO-3 / FSDP+ parametrar84/N + tillfälliga

Med 8 GPU:er och ZeRO-3 blir det ~11 GB per kort i stället för 84.

Effektiv batchstorlek = batch_per_gpu × antal_gpu × accumulation_steps. Den ska hållas konstant när du ändrar uppsättning, annars ändrar du två saker samtidigt och resultatet går inte att jämföra.

Gradient accumulation simulerar stor batch på lite minne: kör k mikrobatchar, ackumulera gradienterna, uppdatera en gång. Matematiskt nästan ekvivalent med en k gånger större batch (skillnaden ligger i batchnorm-statistik, om sådan används).

Vanligaste praktiska felet: att inte skala inlärningstakten när effektiv batch växer, eller att glömma att bara rank 0 ska logga och spara checkpoints — annars skriver alla över varandra.

Kod

# FSDP i PyTorch (motsvarar ZeRO-3)
import torch, torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.utils.data.distributed import DistributedSampler

dist.init_process_group("nccl")
rank, world = dist.get_rank(), dist.get_world_size()
torch.cuda.set_device(rank)

modell = FSDP(modell.cuda(), device_id=rank)
sampler = DistributedSampler(dataset, num_replicas=world, rank=rank, shuffle=True)
loader = torch.utils.data.DataLoader(dataset, batch_size=4, sampler=sampler)

ACC = 4          # effektiv batch = 4 * world * 4
for epok in range(epoker):
    sampler.set_epoch(epok)                   # annars samma ordning varje epok
    for i, (xb, yb) in enumerate(loader):
        loss = kriterium(modell(xb.cuda()), yb.cuda()) / ACC
        loss.backward()
        if (i + 1) % ACC == 0:
            opt.step(); opt.zero_grad()
    if rank == 0:
        spara_checkpoint(modell, epok)        # bara en process skriver

Behärskning innebär

  • Förklarar data-, tensor- och pipelineparallellism
  • Väljer strategi efter modellstorlek och hårdvara
  • Känner till ZeRO/FSDP och gradient accumulation

Logga in för att göra övningarna och bygga upp din behärskning.

Källor

Alla källor och licenser