Gradient checkpointing et SGD local

Entraîner efficacement des modèles d’IA avec PyTorch

Dennis Lee

Data Engineer, Amazon

Améliorer l’efficacité d’entraînement

 

 

Icônes représentant l’efficacité mémoire, l’efficacité de communication et l’efficacité de calcul.

Entraîner efficacement des modèles d’IA avec PyTorch

Le gradient checkpointing améliore l’efficacité mémoire

 

 

Icônes représentant l’efficacité mémoire, l’efficacité de communication et l’efficacité de calcul.

Entraîner efficacement des modèles d’IA avec PyTorch

La SGD locale traite l’efficacité de communication

 

 

Icônes représentant l’efficacité mémoire, l’efficacité de communication et l’efficacité de calcul.

Entraîner efficacement des modèles d’IA avec PyTorch

Qu’est-ce que le gradient checkpointing ?

  • Gradient checkpointing : réduire la mémoire en choisissant quelles activations conserver
  • Exemple : calculer A + B = C

Graphique illustrant le gradient checkpointing avec des nœuds et des arêtes

Entraîner efficacement des modèles d’IA avec PyTorch

Qu’est-ce que le gradient checkpointing ?

  • Gradient checkpointing : réduire la mémoire en choisissant quelles activations conserver
  • Exemple : calculer A + B = C
    • D’abord calculer A, B, puis calculer C

Graphique illustrant le gradient checkpointing avec des nœuds et des arêtes

Entraîner efficacement des modèles d’IA avec PyTorch

Qu’est-ce que le gradient checkpointing ?

  • Gradient checkpointing : réduire la mémoire en choisissant quelles activations conserver
  • Exemple : calculer A + B = C
    • D’abord calculer A, B, puis calculer C
    • A, B ne sont plus utiles pour le reste du passage avant
  • Doit-on conserver ou supprimer A et B ?

Graphique illustrant le gradient checkpointing avec des nœuds et des arêtes

Entraîner efficacement des modèles d’IA avec PyTorch

Qu’est-ce que le gradient checkpointing ?

  • Gradient checkpointing : réduire la mémoire en choisissant quelles activations conserver
  • Exemple : calculer A + B = C
    • D’abord calculer A, B, puis calculer C
    • A, B ne sont plus utiles pour le reste du passage avant
  • Doit-on conserver ou supprimer A et B ?
    • Sans gradient checkpointing : conserver A, B

Graphique illustrant le gradient checkpointing avec des nœuds et des arêtes

Entraîner efficacement des modèles d’IA avec PyTorch

Qu’est-ce que le gradient checkpointing ?

  • Gradient checkpointing : réduire la mémoire en choisissant quelles activations conserver
  • Exemple : calculer A + B = C
    • D’abord calculer A, B, puis calculer C
    • A, B ne sont plus utiles pour le reste du passage avant
  • Doit-on conserver ou supprimer A et B ?
    • Sans gradient checkpointing : conserver A, B
    • Avec gradient checkpointing : supprimer A, B

Graphique illustrant le gradient checkpointing avec des nœuds et des arêtes

Entraîner efficacement des modèles d’IA avec PyTorch

Qu’est-ce que le gradient checkpointing ?

  • Gradient checkpointing : réduire la mémoire en choisissant quelles activations conserver
  • Exemple : calculer A + B = C
    • D’abord calculer A, B, puis calculer C
    • A, B ne sont plus utiles pour le reste du passage avant
  • Doit-on conserver ou supprimer A et B ?
    • Sans gradient checkpointing : conserver A, B
    • Avec gradient checkpointing : supprimer A, B
    • Recalculer A, B pendant la rétropropagation

Graphique illustrant le gradient checkpointing avec des nœuds et des arêtes

Entraîner efficacement des modèles d’IA avec PyTorch

Qu’est-ce que le gradient checkpointing ?

  • Gradient checkpointing : réduire la mémoire en choisissant quelles activations conserver
  • Exemple : calculer A + B = C
    • D’abord calculer A, B, puis calculer C
    • A, B ne sont plus utiles pour le reste du passage avant
  • Doit-on conserver ou supprimer A et B ?
    • Sans gradient checkpointing : conserver A, B
    • Avec gradient checkpointing : supprimer A, B
    • Recalculer A, B pendant la rétropropagation
    • Si B est coûteux à recalculer, le conserver

Graphique illustrant le gradient checkpointing avec des nœuds et des arêtes

Entraîner efficacement des modèles d’IA avec PyTorch

Trainer et Accelerator

Graphique comparant la facilité d’usage et la capacité de personnalisation pour Accelerator et Trainer.

Entraîner efficacement des modèles d’IA avec PyTorch

Trainer et Accelerator

Graphique comparant la facilité d’usage et la capacité de personnalisation pour Accelerator et Trainer.

Entraîner efficacement des modèles d’IA avec PyTorch

Gradient checkpointing avec Trainer

training_args = TrainingArguments(output_dir="./results",
                                  evaluation_strategy="epoch",
                                  gradient_accumulation_steps=4)







Entraîner efficacement des modèles d’IA avec PyTorch

Gradient checkpointing avec Trainer

training_args = TrainingArguments(output_dir="./results",
                                  evaluation_strategy="epoch",
                                  gradient_accumulation_steps=4,
                                  gradient_checkpointing=True)

trainer = Trainer(model=model, args=training_args, train_dataset=dataset["train"], eval_dataset=dataset["validation"], compute_metrics=compute_metrics)
trainer.train()
{'epoch': 1.0, 'eval_loss': 0.73, 'eval_accuracy': 0.03, 'eval_f1': 0.05}
Entraîner efficacement des modèles d’IA avec PyTorch

De Trainer à Accelerator

Graphique comparant la facilité d’usage et la capacité de personnalisation pour Accelerator et Trainer.

Entraîner efficacement des modèles d’IA avec PyTorch

Gradient checkpointing avec Accelerator

accelerator = Accelerator(gradient_accumulation_steps=2)


for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        inputs, targets = batch["input_ids"], batch["labels"]
        outputs = model(inputs, labels=targets)
        loss = outputs.loss
        accelerator.backward(loss)
        optimizer.step()
        lr_scheduler.step()
        optimizer.zero_grad()
Entraîner efficacement des modèles d’IA avec PyTorch

Gradient checkpointing avec Accelerator

accelerator = Accelerator(gradient_accumulation_steps=2)
model.gradient_checkpointing_enable()

for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        inputs, targets = batch["input_ids"], batch["labels"]
        outputs = model(inputs, labels=targets)
        loss = outputs.loss
        accelerator.backward(loss)
        optimizer.step()
        lr_scheduler.step()
        optimizer.zero_grad()
Entraîner efficacement des modèles d’IA avec PyTorch

La SGD locale améliore l’efficacité de communication

 

 

Icônes représentant l’efficacité mémoire, l’efficacité de communication et l’efficacité de calcul.

Entraîner efficacement des modèles d’IA avec PyTorch

Qu’est-ce que la SGD locale ?

Schéma montrant le fonctionnement de la SGD locale en synchronisant les gradients après un certain nombre d’étapes.

  • Chaque dispositif calcule les gradients en parallèle
Entraîner efficacement des modèles d’IA avec PyTorch

Qu’est-ce que la SGD locale ?

Schéma montrant le fonctionnement de la SGD locale en synchronisant les gradients après un certain nombre d’étapes.

  • Chaque dispositif calcule les gradients en parallèle
  • Synchronisation des gradients : le nœud pilote met à jour les paramètres du modèle sur chaque dispositif
  • SGD locale : réduire la fréquence de synchronisation des gradients
Entraîner efficacement des modèles d’IA avec PyTorch

SGD locale avec Accelerator





for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        inputs, targets = batch["input_ids"], batch["labels"]
        outputs = model(inputs, labels=targets)
        loss = outputs.loss
        accelerator.backward(loss)
        optimizer.step()
        lr_scheduler.step()
        optimizer.zero_grad()

Entraîner efficacement des modèles d’IA avec PyTorch

SGD locale avec Accelerator

from accelerate.local_sgd import LocalSGD

with LocalSGD(accelerator=accelerator, model=model, local_sgd_steps=8, 
              enabled=True) as local_sgd:
    for index, batch in enumerate(dataloader):
        with accelerator.accumulate(model):
            inputs, targets = batch["input_ids"], batch["labels"]
            outputs = model(inputs, labels=targets)
            loss = outputs.loss
            accelerator.backward(loss)
            optimizer.step()
            lr_scheduler.step()
            optimizer.zero_grad()

Entraîner efficacement des modèles d’IA avec PyTorch

SGD locale avec Accelerator

from accelerate.local_sgd import LocalSGD

with LocalSGD(accelerator=accelerator, model=model, local_sgd_steps=8, 
              enabled=True) as local_sgd:
    for index, batch in enumerate(dataloader):
        with accelerator.accumulate(model):
            inputs, targets = batch["input_ids"], batch["labels"]
            outputs = model(inputs, labels=targets)
            loss = outputs.loss
            accelerator.backward(loss)
            optimizer.step()
            lr_scheduler.step()
            optimizer.zero_grad()
            local_sgd.step()
Entraîner efficacement des modèles d’IA avec PyTorch

Passons à la pratique !

Entraîner efficacement des modèles d’IA avec PyTorch

Preparing Video For Download...