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 :

  1. L’état de l’optimiseur de ZeRO-1 n’est pas là où on le cherche. ZeroRedundancyOptimizer est une enveloppe : son attribut state reste 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.
  2. 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 explicitement device_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.
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()
'''

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 resultats

Le 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 équivalents float32 : 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_args attend 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, et gloo bloque 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 resultats
HORAIRES = (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 × horizon avant 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.

Retour au sommet