TRAVAILLEUR_COMPTES = r'''
import json
import os
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import ShardingStrategy
from torch.distributed.optim import ZeroRedundancyOptimizer
RANG = int(os.environ["RANK"])
MONDE = int(os.environ["WORLD_SIZE"])
STRATEGIE = os.environ["STRATEGIE"]
COUCHES, LARGEUR, LOT, PAS = 8, 256, 32, 3
def octets(tenseurs):
return sum(t.numel() * t.element_size() for t in tenseurs if torch.is_tensor(t))
def etat_optimiseur(opt):
"""Etat reellement detenu par ce rang, y compris pour ZeroRedundancyOptimizer."""
total = 0
for source in (getattr(opt, "optim", None), opt):
if source is None:
continue
total = max(total, octets([v for etat in source.state.values()
for v in etat.values()]))
return total
def modele(graine=0):
torch.manual_seed(graine)
return nn.Sequential(*[nn.Linear(LARGEUR, LARGEUR) for _ in range(COUCHES)])
def entraine(m, opt, journal):
for _ in range(PAS):
x = torch.randn(LOT, LARGEUR)
opt.zero_grad(set_to_none=True)
y = m(x)
if journal.get("apres_forward") is None:
journal["apres_forward"] = octets(list(m.parameters()))
y.sum().backward()
if journal.get("apres_backward") is None:
journal["apres_backward"] = octets(list(m.parameters()))
opt.step()
def main():
dist.init_process_group("gloo", init_method=os.environ["DIST_INIT"],
rank=RANG, world_size=MONDE)
base = modele()
journal = {"apres_forward": None, "apres_backward": None}
if STRATEGIE == "ddp":
m = nn.parallel.DistributedDataParallel(base)
opt = torch.optim.AdamW(m.parameters(), lr=1e-3)
elif STRATEGIE == "zero1":
m = base
opt = ZeroRedundancyOptimizer(m.parameters(),
optimizer_class=torch.optim.AdamW, lr=1e-3)
elif STRATEGIE in ("zero2", "zero3"):
strategie = (ShardingStrategy.SHARD_GRAD_OP if STRATEGIE == "zero2"
else ShardingStrategy.FULL_SHARD)
m = FSDP(base, sharding_strategy=strategie,
device_id=torch.device("cpu"))
opt = torch.optim.AdamW(m.parameters(), lr=1e-3)
else:
raise SystemExit("strategie inconnue : " + STRATEGIE)
entraine(m, opt, journal)
poids = octets(list(m.parameters()))
grads = octets([p.grad for p in m.parameters()])
optim = etat_optimiseur(opt)
sortie = {"rang": RANG, "monde": MONDE, "strategie": STRATEGIE,
"poids": poids, "gradients": grads, "optimiseur": optim,
"total": poids + grads + optim,
"apres_forward": journal["apres_forward"],
"apres_backward": journal["apres_backward"]}
print("RESULTAT " + json.dumps(sortie))
dist.destroy_process_group()
if __name__ == "__main__":
main()
'''3.13 — Découper le modèle : DDP, ZeRO, FSDP — et ce que le pipeline change
Navigation : << 3.12 — Les collectives · 3.13b — Deux vraies cartes : nccl >> · Feuille de route de la série
Le carnet précédent a isolé la brique de communication : un all_reduce en anneau, comparé à celui de la bibliothèque. Ici, on monte d’un étage et on change de question. Un modèle qui ne tient pas sur une carte ne devient pas entraînable parce que la communication est rapide — il le devient parce que ce que chaque rang garde en mémoire diminue. Or « découper » ne veut pas dire la même chose selon la stratégie : DDP réplique tout, ZeRO-1 ne découpe que l’état de l’optimiseur, ZeRO-2 et ZeRO-3 découpent aussi les gradients et les paramètres, et le parallélisme de pipeline découpe le modèle par étages en gardant chaque étage entier.
Ce carnet mesure ces différences au lieu de les réciter. Toutes les mesures sont prises ici, sur CPU, avec les vrais objets de torch.distributed : DistributedDataParallel, ZeroRedundancyOptimizer, FullyShardedDataParallel, et torch.distributed.pipelining.
Ce que le carnet mesure, et sur quoi il se prononce
| Question | Instrument | Ce qui est comparé |
|---|---|---|
| Combien d’octets par rang ? | somme des tenseurs vivants (poids, gradients, état d’optimiseur) | quatre stratégies, à monde 2 et 4 |
| L’état de l’optimiseur est-il vraiment découpé ? | lecture de l’optimiseur local, pas de l’enveloppe | ZeRO-1 contre DDP |
| ZeRO-2 et ZeRO-3 diffèrent-ils au repos ? | mêmes sommes, à deux moments du pas | ce que l’instrument voit, et ce qu’il ne voit pas |
| Combien d’activations en vol ? | torch.autograd.graph.saved_tensors_hooks |
GPipe contre 1F1B, à horizon croissant |
Une remarque de méthode, valable pour tout le carnet : un chiffre de mémoire n’est pas une propriété du modèle, c’est une propriété du couple (stratégie, nombre de rangs). Le même réseau de 526 336 paramètres passe de 8,4 Mo par rang à 2,1 Mo selon ce qu’on découpe — et c’est tout l’objet de ce qui suit.
Le dispositif : quatre stratégies, un seul modèle
Le modèle de travail est volontairement petit — huit couches Linear(256, 256) — pour que la mesure soit lisible : on veut voir les rapports exacts, pas des ordres de grandeur noyés dans le bruit. Il compte \(P = 8 \times (256 \times 256 + 256) = 526\,336\) paramètres, soit 2 105 344 octets en float32.
Quatre quantités sont intéressantes par rang :
| Quantité | Ce qu’elle contient | Taille pleine |
|---|---|---|
| poids | les paramètres du modèle | \(4P\) |
| gradients | un gradient par paramètre | \(4P\) |
| état d’optimiseur | pour AdamW, deux tenseurs par paramètre (exp_avg, exp_avg_sq) |
\(8P\) |
Chaque rang est un processus séparé, lancé par le carnet, qui se donne rendez-vous sur un store fichier — le même dispositif que le carnet 3.12, et pour la même raison : c’est le seul moyen d’obtenir de vrais rangs sur une machine sans ordonnanceur de cluster.
Deux pièges de mesure, tous deux rencontrés pour de vrai en écrivant ce carnet :
- L’état de l’optimiseur de ZeRO-1 n’est pas là où on le cherche.
ZeroRedundancyOptimizerest une enveloppe : son attributstatereste vide, l’état réel vit dans l’optimiseur local qu’elle contient (opt.optim.state). Une lecture naïve rend zéro octet et fait croire à une économie totale — un chiffre faux qui a l’air d’un excellent résultat. - FSDP refuse le CPU par défaut :
FSDP needs a non-CPU accelerator device. Ce refus ne tient que dans la branche d’auto-détection ; passer explicitementdevice_id=torch.device("cpu")le contourne légitimement. C’est ce que fait le travailleur ci-dessous, et c’est ce qui rend la comparaison possible ici.
Le travailleur, et son lanceur
Le travailleur reçoit sa stratégie et son rang par l’environnement, s’initialise sur gloo, construit le modèle, l’enveloppe selon la stratégie demandée, puis exécute trois pas d’entraînement. Il ne rend pas une impression mais une ligne RESULTAT en JSON, que le lanceur ramasse.
import json
import os
import pathlib
import subprocess
import sys
import tempfile
def lance(monde, travailleur, strategie, delai=240):
"""Lance `monde` processus du travailleur et ramasse leurs lignes RESULTAT."""
dossier = pathlib.Path(tempfile.mkdtemp(prefix="comptes_"))
chemin = dossier / "travailleur.py"
chemin.write_text(travailleur, encoding="utf-8")
store = pathlib.Path(tempfile.mkdtemp(prefix="comptes_store_"))
init = "file:///" + str(store / "rv").replace("\\", "/")
env = dict(os.environ, DIST_INIT=init, WORLD_SIZE=str(monde), STRATEGIE=strategie)
procs = [subprocess.Popen([sys.executable, str(chemin)],
env=dict(env, RANK=str(r)),
stdout=subprocess.PIPE, stderr=subprocess.PIPE,
text=True, encoding="utf-8") for r in range(monde)]
resultats, pannes = [], []
for rang, proc in enumerate(procs):
try:
sortie, erreur = proc.communicate(timeout=delai)
except subprocess.TimeoutExpired:
proc.kill()
sortie, erreur = proc.communicate()
pannes.append("rang %d : blocage apres %d s" % (rang, delai))
continue
if proc.returncode != 0:
pannes.append("rang %d : code %d\n%s"
% (rang, proc.returncode, erreur.strip()[-600:]))
for ligne in sortie.splitlines():
if ligne.startswith("RESULTAT "):
resultats.append(json.loads(ligne[len("RESULTAT "):]))
if pannes:
raise RuntimeError("echec de %s a monde %d :\n%s"
% (strategie, monde, "\n".join(pannes)))
return resultatsLe lanceur
Une fonction, quatre stratégies, deux tailles de monde. Le délai de garde n’est pas décoratif : un rang qui ne rejoint pas le groupe laisse les autres bloqués, et sans délai le carnet attendrait indéfiniment.
STRATEGIES = ("ddp", "zero1", "zero2", "zero3")
MONDES = (2, 4)
mesures = {}
for monde in MONDES:
for strategie in STRATEGIES:
rangs = lance(monde, TRAVAILLEUR_COMPTES, strategie)
totaux = sorted({r["total"] for r in rangs})
print("monde=%d %-6s rangs=%d totaux=%s"
% (monde, strategie, len(rangs), totaux))
mesures[(monde, strategie)] = rangs[0]
print()
print("relevé retenu (rang 0) :")
for (monde, strategie), m in mesures.items():
print(" monde=%d %-6s poids=%9d grads=%9d optim=%9d total=%9d"
% (monde, strategie, m["poids"], m["gradients"], m["optimiseur"], m["total"]))monde=2 ddp rangs=2 totaux=[8421440]
monde=2 zero1 rangs=2 totaux=[6316064]
monde=2 zero2 rangs=2 totaux=[4210692]
monde=2 zero3 rangs=2 totaux=[4210692]
monde=4 ddp rangs=4 totaux=[8421440]
monde=4 zero1 rangs=4 totaux=[5263376]
monde=4 zero2 rangs=4 totaux=[2105348]
monde=4 zero3 rangs=4 totaux=[2105348]
relevé retenu (rang 0) :
monde=2 ddp poids= 2105344 grads= 2105344 optim= 4210752 total= 8421440
monde=2 zero1 poids= 2105344 grads= 2105344 optim= 2105376 total= 6316064
monde=2 zero2 poids= 1052672 grads= 1052672 optim= 2105348 total= 4210692
monde=2 zero3 poids= 1052672 grads= 1052672 optim= 2105348 total= 4210692
monde=4 ddp poids= 2105344 grads= 2105344 optim= 4210752 total= 8421440
monde=4 zero1 poids= 2105344 grads= 2105344 optim= 1052688 total= 5263376
monde=4 zero2 poids= 526336 grads= 526336 optim= 1052676 total= 2105348
monde=4 zero3 poids= 526336 grads= 526336 optim= 1052676 total= 2105348
P = 8 * (256 * 256 + 256)
print("%-8s %-6s %10s %10s %10s %10s %s"
% ("monde", "strat", "poids", "grads", "optim", "total", "en multiples de P"))
for (monde, strategie), m in mesures.items():
print("%-8d %-6s %10d %10d %10d %10d %.2f P"
% (monde, strategie, m["poids"], m["gradients"], m["optimiseur"],
m["total"], m["total"] / (4 * P)))monde strat poids grads optim total en multiples de P
2 ddp 2105344 2105344 4210752 8421440 4.00 P
2 zero1 2105344 2105344 2105376 6316064 3.00 P
2 zero2 1052672 1052672 2105348 4210692 2.00 P
2 zero3 1052672 1052672 2105348 4210692 2.00 P
4 ddp 2105344 2105344 4210752 8421440 4.00 P
4 zero1 2105344 2105344 1052688 5263376 2.50 P
4 zero2 526336 526336 1052676 2105348 1.00 P
4 zero3 526336 526336 1052676 2105348 1.00 P
Premier dépouillement : ce que chaque stratégie découpe
Les huit configurations sont lancées, et chaque rang rend sa ligne. On regroupe par (monde, stratégie) et on lit la première ligne de chaque groupe — tous les rangs d’une même configuration doivent rendre le même total, sinon la mesure est fausse.
ZeRO-1 ne découpe que l’état de l’optimiseur
À monde 2, ZeRO-1 divise l’état d’optimiseur par deux et laisse poids et gradients intacts ; à monde 4, par quatre. C’est exactement ce qu’annonce son nom : le premier étage de ZeRO ne touche qu’à l’état de l’optimiseur, qui est la plus grosse des trois quantités (8P contre 4P et 4P).
Le tableau le chiffre : DDP vaut 16P quel que soit le monde, ZeRO-1 vaut \(8P + 8P/\text{monde}\), et c’est bien ce qu’on observe — 6 316 064 octets à monde 2 (12P), 5 263 376 à monde 4 (10P).
Deux échelles de P cohabitent ici, sans erreur. Dans les formules ci-dessus, « 16P » se lit seize octets par paramètre (poids bf16 2 + gradients bf16 2 + état Adam fp32 12) : 16 × 526 336 = 8 421 376 octets, le total DDP mesuré. Le tableau de la cellule suivante, lui, divise chaque total par 4 × P — quatre octets par paramètre
float32— et affiche donc « 4.00 P » sur la même ligne DDP. La prose compte en octets par paramètre (l’échelle des formules ZeRO de la littérature), le tableau en jeux de paramètres équivalentsfloat32: même mémoire, deux unités, un facteur 4.
À quoi sert de ne découper « que » l’optimiseur ? Parce que c’est là que va la mémoire quand on entraîne en précision mixte : avec des poids en bfloat16 et un optimiseur en float32, l’état d’optimiseur pèse plusieurs fois le modèle. Le premier étage de ZeRO attaque le poste le plus lourd sans toucher au chemin critique de la communication des gradients.
ZeRO-2 et ZeRO-3 : même état au repos, et c’est un vrai résultat
Le tableau montre une chose qui surprend la première fois : ZeRO-2 et ZeRO-3 rendent des totaux identiques, au repos, aux deux tailles de monde. Ce n’est pas une erreur de mesure, et ce n’est pas non plus une erreur de la bibliothèque.
La documentation de ShardingStrategy dit pourquoi : SHARD_GRAD_OP (ZeRO-2) découpe les gradients et l’état d’optimiseur pendant le calcul, et « additionally, parameters are sharded outside computation ». FULL_SHARD (ZeRO-3) découpe les mêmes trois quantités, mais reshard les paramètres juste après le forward, là où SHARD_GRAD_OP les garde déployés jusqu’à la fin du backward.
Autrement dit : les deux stratégies laissent le même repos, et diffèrent par le pic pendant le backward. Un instrument qui somme l’état au repos — celui de ce carnet — ne peut structurellement pas les distinguer, et il faut le dire plutôt que de fabriquer une différence. Ce que SHARD_GRAD_OP achète, c’est un all-gather de moins avant le backward, au prix de paramètres déployés plus longtemps : un arbitrage de temps contre pic de mémoire, invisible sur un relevé au repos.
C’est le genre d’endroit où un carnet pédagogique peut mentir sans le vouloir : deux colonnes égales invitent à écrire « ZeRO-2 et ZeRO-3 sont équivalents », ce qui est faux. Elles sont équivalentes au repos, et c’est tout ce que cette mesure établit.
Le pipeline : découper par étages, et l’ordre des micro-lots
Le parallélisme de données découpe les données ; le pipeline découpe le modèle en étages, un par rang, et fait circuler les activations de l’un à l’autre. Un lot est alors découpé en micro-lots qui traversent les étages à la queue leu leu — c’est ce qui permet aux étages de travailler en parallèle au lieu de s’attendre.
Reste à choisir l’ordre : soit on fait traverser tous les micro-lots en avant avant de commencer le moindre backward (GPipe), soit on entrelace — dès qu’un micro-lot a fini sa course avant, on lance son backward (1F1B). L’ordre ne change pas le résultat mathématique, mais il change ce qui est vivant en mémoire à un instant donné.
Pour le mesurer, il faut un instrument qui survive au tracé du graphe. torch.distributed.pipelining construit un GraphModule par tracé fx : patcher la méthode forward d’une couche ne compte alors que les forwards du tracé, pas ceux de l’exécution — un piège qui rend des compteurs identiques partout. L’instrument retenu est torch.autograd.graph.saved_tensors_hooks : pack est appelé quand un tenseur est retenu pour le backward, unpack quand il est consommé. Le pic de tenseurs simultanément retenus est l’occupation mémoire des activations.
Deux règles de l’API, apprises à la dure et qui bloquent sans message si on les ignore :
mb_argsattend un exemple de la taille d’un micro-lot, pas du lot complet. Avec le lot complet, le rang destinataire attend une forme que l’émetteur n’envoie pas, etgloobloque indéfiniment.- les cibles se passent au lot complet et au dernier étage — le planificateur les découpe lui-même.
TRAVAILLEUR_PIPELINE = r'''
import json
import os
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.distributed.pipelining import (Schedule1F1B, ScheduleGPipe, SplitPoint,
pipeline)
RANG = int(os.environ["RANK"])
MONDE = int(os.environ["WORLD_SIZE"])
HORAIRES = int(os.environ["HORAIRES"])
ORDONNAGEUR = os.environ["ORDONNAGEUR"]
LARGEUR, LOT = 32, 8
class Bloc(nn.Module):
def __init__(self, largeur):
super().__init__()
self.couche = nn.Linear(largeur, largeur)
def forward(self, x):
return torch.relu(self.couche(x))
class Reseau(nn.Module):
def __init__(self, etages, largeur):
super().__init__()
self.etages = nn.ModuleList([Bloc(largeur) for _ in range(etages)])
def forward(self, x):
for etage in self.etages:
x = etage(x)
return x
def main():
dist.init_process_group("gloo", init_method=os.environ["DIST_INIT"],
rank=RANG, world_size=MONDE)
torch.manual_seed(0)
reseau = Reseau(MONDE, LARGEUR)
exemple = torch.randn(LOT, LARGEUR)
micro = torch.randn(LOT // HORAIRES, LARGEUR)
spec = {"etages.%d" % i: SplitPoint.BEGINNING for i in range(1, MONDE)}
pipe = pipeline(reseau, mb_args=(micro,), split_spec=spec)
stage = pipe.build_stage(RANG, torch.device("cpu"))
classe = ScheduleGPipe if ORDONNAGEUR == "gpipe" else Schedule1F1B
ordonnanceur = classe(stage, HORAIRES, loss_fn=torch.nn.functional.mse_loss)
journal = {"vivants": 0, "pic": 0, "sauvegardes": 0, "consommes": 0,
"octets_vivants": 0, "pic_octets": 0, "sauvegardes_avant_1er_unpack": None}
def pack(t):
journal["vivants"] += 1
journal["sauvegardes"] += 1
journal["octets_vivants"] += t.numel() * t.element_size()
if journal["vivants"] > journal["pic"]:
journal["pic"] = journal["vivants"]
if journal["octets_vivants"] > journal["pic_octets"]:
journal["pic_octets"] = journal["octets_vivants"]
return t
def unpack(t):
journal["vivants"] -= 1
journal["octets_vivants"] -= t.numel() * t.element_size()
journal["consommes"] += 1
if journal["sauvegardes_avant_1er_unpack"] is None:
journal["sauvegardes_avant_1er_unpack"] = journal["sauvegardes"]
return t
cibles = torch.ones(LOT, LARGEUR)
with torch.autograd.graph.saved_tensors_hooks(pack, unpack):
if stage.is_first:
ordonnanceur.step(exemple)
elif stage.is_last:
ordonnanceur.step(target=cibles)
else:
ordonnanceur.step()
resultat = {"rang": RANG, "monde": MONDE, "ordonnanceur": ORDONNAGEUR,
"horaires": HORAIRES, "premier_etage": stage.is_first,
"dernier_etage": stage.is_last,
"pic_vivants": journal["pic"], "pic_octets": journal["pic_octets"],
"sauvegardes": journal["sauvegardes"],
"consommes": journal["consommes"],
"sauvegardes_avant_1er_unpack": journal["sauvegardes_avant_1er_unpack"]}
print("RESULTAT " + json.dumps(resultat))
dist.destroy_process_group()
if __name__ == "__main__":
main()
'''Le travailleur de pipeline, et son lanceur
Le travailleur instrumente les sauvegardes de tenseurs, construit le pipeline avec un découpage explicite (SplitPoint.BEGINNING sur chaque étage sauf le premier), puis exécute un pas du planificateur demandé.
def lance_pipeline(monde, travailleur, ordonnageur, horaires, delai=150):
"""Meme dispositif que `lance`, avec l'ordonnanceur et l'horizon en plus."""
dossier = pathlib.Path(tempfile.mkdtemp(prefix="pipe_"))
chemin = dossier / "travailleur.py"
chemin.write_text(travailleur, encoding="utf-8")
store = pathlib.Path(tempfile.mkdtemp(prefix="pipe_store_"))
init = "file:///" + str(store / "rv").replace("\\", "/")
env = dict(os.environ, DIST_INIT=init, WORLD_SIZE=str(monde),
ORDONNAGEUR=ordonnageur, HORAIRES=str(horaires))
procs = [subprocess.Popen([sys.executable, str(chemin)],
env=dict(env, RANK=str(r)),
stdout=subprocess.PIPE, stderr=subprocess.PIPE,
text=True, encoding="utf-8") for r in range(monde)]
resultats, pannes = [], []
for rang, proc in enumerate(procs):
try:
sortie, erreur = proc.communicate(timeout=delai)
except subprocess.TimeoutExpired:
proc.kill()
sortie, erreur = proc.communicate()
pannes.append("rang %d : blocage apres %d s" % (rang, delai))
continue
if proc.returncode != 0:
pannes.append("rang %d : code %d\n%s"
% (rang, proc.returncode, erreur.strip()[-600:]))
for ligne in sortie.splitlines():
if ligne.startswith("RESULTAT "):
resultats.append(json.loads(ligne[len("RESULTAT "):]))
if pannes:
raise RuntimeError("echec de %s horizon %d :\n%s"
% (ordonnageur, horaires, "\n".join(pannes)))
return resultatsHORAIRES = (2, 4, 8)
mesures_pipeline = {}
for ordonnageur in ("gpipe", "1f1b"):
for horizon in HORAIRES:
rangs = lance_pipeline(2, TRAVAILLEUR_PIPELINE, ordonnageur, horizon)
dernier = [r for r in rangs if r["dernier_etage"]][0]
mesures_pipeline[(ordonnageur, horizon)] = dernier
print("%-6s horizon=%d pic du dernier etage : %2d tenseurs vivants, "
"%d sauvegardes avant le premier unpack"
% (ordonnageur, horizon, dernier["pic_vivants"],
dernier["sauvegardes_avant_1er_unpack"]))gpipe horizon=2 pic du dernier etage : 10 tenseurs vivants, 10 sauvegardes avant le premier unpack
gpipe horizon=4 pic du dernier etage : 20 tenseurs vivants, 20 sauvegardes avant le premier unpack
gpipe horizon=8 pic du dernier etage : 40 tenseurs vivants, 40 sauvegardes avant le premier unpack
1f1b horizon=2 pic du dernier etage : 5 tenseurs vivants, 5 sauvegardes avant le premier unpack
1f1b horizon=4 pic du dernier etage : 5 tenseurs vivants, 5 sauvegardes avant le premier unpack
1f1b horizon=8 pic du dernier etage : 5 tenseurs vivants, 5 sauvegardes avant le premier unpack
Le relevé : deux ordonnancements, trois horizons
Six configurations, mesurées sur deux rangs. Pour chacune, on ne retient que le dernier étage : c’est lui qui porte la perte, donc lui qui retient les activations et leurs cibles, et c’est là que l’écart entre les deux ordres se voit.
print("%-8s %8s %14s %16s %12s"
% ("ordonn.", "horizon", "pic vivants", "sauv. avant 1er", "sauv. totales"))
for (ordonnageur, horizon), m in mesures_pipeline.items():
print("%-8s %8d %14d %16d %12d"
% (ordonnageur, horizon, m["pic_vivants"],
m["sauvegardes_avant_1er_unpack"], m["sauvegardes"]))ordonn. horizon pic vivants sauv. avant 1er sauv. totales
gpipe 2 10 10 10
gpipe 4 20 20 20
gpipe 8 40 40 40
1f1b 2 5 5 10
1f1b 4 5 5 20
1f1b 8 5 5 40
Lecture : le pic croît avec l’horizon, ou pas
Le contraste est net, et il est exactement celui qu’annonce la théorie :
- GPipe : le dernier étage retient 5 tenseurs par micro-lot et en a retenu
5 × horizonavant que le premier backward ne consomme quoi que ce soit. À horizon 8, cela fait 40 tenseurs vivants simultanément — la mémoire des activations croît linéairement avec le nombre de micro-lots. - 1F1B : le premier backward part dès le premier micro-lot, et le pic reste constant à 5 tenseurs, quel que soit l’horizon. La colonne « sauvegardes avant le premier unpack » le montre directement : 5 pour 1F1B contre 40 pour GPipe à horizon 8, alors que les deux font bien 40 sauvegardes au total.
C’est l’arbitrage du parallélisme de pipeline : GPipe est plus simple et laisse les étages travailler sans se chevaucher, 1F1B borne la mémoire mais fait cohabiter forwards et backwards. Le nombre de micro-lots n’est donc pas un réglage neutre — c’est le paramètre qui décide si le pipeline tient sur les cartes ou non.
Ce que cet instrument voit, et ce qu’il ne voit pas
Il compte les tenseurs d’activation retenus pour le backward. Il ne voit pas les tampons de communication entre étages, ni les paramètres déployés temporairement par FSDP, ni la fragmentation de l’allocateur. Un pic de 40 tenseurs n’est donc pas « 40 fois plus de mémoire » : c’est 40 tenseurs de tailles différentes, et ce carnet mesure le nombre, pas les octets — précisément parce que la taille d’un micro-lot change avec l’horizon (le lot est découpé), ce qui rendrait la comparaison en octets trompeuse.
Exercices
Les trois exercices prolongent la mesure : on y prédit avant de vérifier, parce que c’est la seule façon de savoir si l’on a compris ce que les chiffres racontent.
# Exercice 1 -- Predire les octets par rang, puis verifier contre la mesure.
# TODO etudiant : completer octets_par_rang(strategie, monde) qui rend le total
# par rang EN OCTETS, a partir des formules : P = 526336 parametres,
# soit 4*P pour les poids, 4*P pour les gradients, 8*P pour AdamW.
# Indice : ddp replique tout ; zero1 ne decoupe QUE l'etat d'optimiseur ;
# zero2 et zero3 decoupent les trois quantites par le monde.
# Critere de reussite : l'ecart relatif avec la mesure doit rester sous 1 %
# pour les huit configurations (les quelques octets d'ecart sont des
# constantes par rang, pas une erreur de formule).
P = 8 * (256 * 256 + 256)
def octets_par_rang(strategie, monde):
return None
print("Exercice a completer")Exercice a completer
# Exercice 2 -- L'horizon du pipeline.
# TODO etudiant : lancer le travailleur de pipeline a horizon 6 pour les deux
# ordonnancements, et reporter le pic du DERNIER etage.
# Indice : le pic du dernier etage vaut 5 * horizon pour gpipe, et reste a 5
# pour 1f1b -- la loi se lit sur les trois horizons du tableau.
# Critere de reussite : pic(gpipe, 6) == 30 et pic(1f1b, 6) == 5.
# Attention : la taille d'un micro-lot vaut LOT // horizon ; a horizon 6 avec
# LOT = 8, le decoupage n'est pas exact -- choisir un horizon qui
# divise le lot, ou ajuster LOT.
mesures_h6 = None # TODO etudiant : remplacer par les lancements
print("Exercice a completer")Exercice a completer
# Exercice 3 -- Le total sur la machine entiere, pas par rang.
# TODO etudiant : completer total_machine(strategie, monde), qui multiplie le
# total par rang par le nombre de rangs, et l'appeler pour monde dans
# (1, 2, 4, 8).
# Indice : c'est ici que l'argument du decoupage apparait -- comparer les
# colonnes, pas les lignes.
# Critere de reussite : le total machine de ddp croit lineairement avec le
# monde, celui de zero3 reste constant a 16*P.
def total_machine(strategie, monde):
return None
print("Exercice a completer")Exercice a completer
Suite
Le carnet a mesuré ce que chaque stratégie garde par rang, et ce que l’ordre des micro-lots change à l’occupation. Les deux moitiés se rejoignent sur un même point : ce qui limite un entraînement distribué n’est presque jamais le calcul, c’est ce qui reste vivant en mémoire à l’instant le plus défavorable — et cet instant n’est pas le même selon qu’on découpe les données, les états, le modèle, ou la suite des micro-lots.