Gradient checkpointing en lokale SGD

Efficiënt AI-modellen trainen met PyTorch

Dennis Lee

Data Engineer, Amazon

Training efficiënter maken

 

 

Pictogrammen voor geheugenefficiëntie, communicatiesnelheid en rekenefficiëntie.

Efficiënt AI-modellen trainen met PyTorch

Gradient checkpointing verbetert geheugenefficiëntie

 

 

Pictogrammen voor geheugenefficiëntie, communicatiesnelheid en rekenefficiëntie.

Efficiënt AI-modellen trainen met PyTorch

Lokale SGD pakt communicatie-efficiëntie aan

 

 

Pictogrammen voor geheugenefficiëntie, communicatiesnelheid en rekenefficiëntie.

Efficiënt AI-modellen trainen met PyTorch

Wat is gradient checkpointing?

  • Gradient checkpointing: minder geheugen door te kiezen welke activaties je bewaart
  • Voorbeeld: bereken A + B = C

Grafiek die gradient checkpointing toont met knooppunten en randen

Efficiënt AI-modellen trainen met PyTorch

Wat is gradient checkpointing?

  • Gradient checkpointing: minder geheugen door te kiezen welke activaties je bewaart
  • Voorbeeld: bereken A + B = C
    • Eerst A en B, daarna C

Grafiek die gradient checkpointing toont met knooppunten en randen

Efficiënt AI-modellen trainen met PyTorch

Wat is gradient checkpointing?

  • Gradient checkpointing: minder geheugen door te kiezen welke activaties je bewaart
  • Voorbeeld: bereken A + B = C
    • Eerst A en B, daarna C
    • A en B zijn niet nodig voor de rest van de forward pass
  • Moeten we A en B bewaren of weggooien?

Grafiek die gradient checkpointing toont met knooppunten en randen

Efficiënt AI-modellen trainen met PyTorch

Wat is gradient checkpointing?

  • Gradient checkpointing: minder geheugen door te kiezen welke activaties je bewaart
  • Voorbeeld: bereken A + B = C
    • Eerst A en B, daarna C
    • A en B zijn niet nodig voor de rest van de forward pass
  • Moeten we A en B bewaren of weggooien?
    • Zonder gradient checkpointing: bewaar A en B

Grafiek die gradient checkpointing toont met knooppunten en randen

Efficiënt AI-modellen trainen met PyTorch

Wat is gradient checkpointing?

  • Gradient checkpointing: minder geheugen door te kiezen welke activaties je bewaart
  • Voorbeeld: bereken A + B = C
    • Eerst A en B, daarna C
    • A en B zijn niet nodig voor de rest van de forward pass
  • Moeten we A en B bewaren of weggooien?
    • Zonder gradient checkpointing: bewaar A en B
    • Met gradient checkpointing: verwijder A en B

Grafiek die gradient checkpointing toont met knooppunten en randen

Efficiënt AI-modellen trainen met PyTorch

Wat is gradient checkpointing?

  • Gradient checkpointing: minder geheugen door te kiezen welke activaties je bewaart
  • Voorbeeld: bereken A + B = C
    • Eerst A en B, daarna C
    • A en B zijn niet nodig voor de rest van de forward pass
  • Moeten we A en B bewaren of weggooien?
    • Zonder gradient checkpointing: bewaar A en B
    • Met gradient checkpointing: verwijder A en B
    • Bereken A en B opnieuw tijdens de backward pass

Grafiek die gradient checkpointing toont met knooppunten en randen

Efficiënt AI-modellen trainen met PyTorch

Wat is gradient checkpointing?

  • Gradient checkpointing: minder geheugen door te kiezen welke activaties je bewaart
  • Voorbeeld: bereken A + B = C
    • Eerst A en B, daarna C
    • A en B zijn niet nodig voor de rest van de forward pass
  • Moeten we A en B bewaren of weggooien?
    • Zonder gradient checkpointing: bewaar A en B
    • Met gradient checkpointing: verwijder A en B
    • Bereken A en B opnieuw tijdens de backward pass
    • Is B duur om te herberekenen? Bewaar B

Grafiek die gradient checkpointing toont met knooppunten en randen

Efficiënt AI-modellen trainen met PyTorch

Trainer en Accelerator

Grafiek die gebruiksgemak versus aanpasbaarheid vergelijkt voor Accelerator en Trainer.

Efficiënt AI-modellen trainen met PyTorch

Trainer en Accelerator

Grafiek die gebruiksgemak versus aanpasbaarheid vergelijkt voor Accelerator en Trainer.

Efficiënt AI-modellen trainen met PyTorch

Gradient checkpointing met Trainer

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







Efficiënt AI-modellen trainen met PyTorch

Gradient checkpointing met 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}
Efficiënt AI-modellen trainen met PyTorch

Van Trainer naar Accelerator

Grafiek die gebruiksgemak versus aanpasbaarheid vergelijkt voor Accelerator en Trainer.

Efficiënt AI-modellen trainen met PyTorch

Gradient checkpointing met 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()
Efficiënt AI-modellen trainen met PyTorch

Gradient checkpointing met 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()
Efficiënt AI-modellen trainen met PyTorch

Lokale SGD verbetert communicatie-efficiëntie

 

 

Pictogrammen voor geheugenefficiëntie, communicatiesnelheid en rekenefficiëntie.

Efficiënt AI-modellen trainen met PyTorch

Wat is lokale SGD?

Diagram dat laat zien hoe lokale SGD werkt door na een aantal stappen te synchroniseren.

  • Elk device berekent parallel de gradiënten
Efficiënt AI-modellen trainen met PyTorch

Wat is lokale SGD?

Diagram dat laat zien hoe lokale SGD werkt door na een aantal stappen te synchroniseren.

  • Elk device berekent parallel de gradiënten
  • Gradiëntsynchronisatie: drivernode werkt modelparameters op elk device bij
  • Lokale SGD: verlaag de frequentie van gradiëntsynchronisatie
Efficiënt AI-modellen trainen met PyTorch

Lokale SGD met 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()

Efficiënt AI-modellen trainen met PyTorch

Lokale SGD met 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()

Efficiënt AI-modellen trainen met PyTorch

Lokale SGD met 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()
Efficiënt AI-modellen trainen met PyTorch

Laten we oefenen!

Efficiënt AI-modellen trainen met PyTorch

Preparing Video For Download...