Gradient checkpointing și SGD local

Antrenament eficient al modelelor AI cu PyTorch

Dennis Lee

Data Engineer, Amazon

Îmbunătățirea eficienței antrenamentului

 

 

Pictograme reprezentând eficiența memoriei, comunicării și calculului.

Antrenament eficient al modelelor AI cu PyTorch

Gradient checkpointing îmbunătățește eficiența memoriei

 

 

Pictograme reprezentând eficiența memoriei, comunicării și calculului.

Antrenament eficient al modelelor AI cu PyTorch

SGD local îmbunătățește eficiența comunicării

 

 

Pictograme reprezentând eficiența memoriei, comunicării și calculului.

Antrenament eficient al modelelor AI cu PyTorch

Ce este gradient checkpointing?

  • Gradient checkpointing: reduce memoria selectând activările de salvat
  • Exemplu: calculăm A + B = C

Grafic ilustrând gradient checkpointing cu noduri și muchii

Antrenament eficient al modelelor AI cu PyTorch

Ce este gradient checkpointing?

  • Gradient checkpointing: reduce memoria selectând activările de salvat
  • Exemplu: calculăm A + B = C
    • Mai întâi calculăm A, B, apoi C

Grafic ilustrând gradient checkpointing cu noduri și muchii

Antrenament eficient al modelelor AI cu PyTorch

Ce este gradient checkpointing?

  • Gradient checkpointing: reduce memoria selectând activările de salvat
  • Exemplu: calculăm A + B = C
    • Mai întâi calculăm A, B, apoi C
    • A, B nu sunt necesare în restul pasului înainte
  • Salvăm sau eliminăm A și B?

Grafic ilustrând gradient checkpointing cu noduri și muchii

Antrenament eficient al modelelor AI cu PyTorch

Ce este gradient checkpointing?

  • Gradient checkpointing: reduce memoria selectând activările de salvat
  • Exemplu: calculăm A + B = C
    • Mai întâi calculăm A, B, apoi C
    • A, B nu sunt necesare în restul pasului înainte
  • Salvăm sau eliminăm A și B?
    • Fără gradient checkpointing: salvăm A, B

Grafic ilustrând gradient checkpointing cu noduri și muchii

Antrenament eficient al modelelor AI cu PyTorch

Ce este gradient checkpointing?

  • Gradient checkpointing: reduce memoria selectând activările de salvat
  • Exemplu: calculăm A + B = C
    • Mai întâi calculăm A, B, apoi C
    • A, B nu sunt necesare în restul pasului înainte
  • Salvăm sau eliminăm A și B?
    • Fără gradient checkpointing: salvăm A, B
    • Gradient checkpointing: eliminăm A, B

Grafic ilustrând gradient checkpointing cu noduri și muchii

Antrenament eficient al modelelor AI cu PyTorch

Ce este gradient checkpointing?

  • Gradient checkpointing: reduce memoria selectând activările de salvat
  • Exemplu: calculăm A + B = C
    • Mai întâi calculăm A, B, apoi C
    • A, B nu sunt necesare în restul pasului înainte
  • Salvăm sau eliminăm A și B?
    • Fără gradient checkpointing: salvăm A, B
    • Gradient checkpointing: eliminăm A, B
    • Recalculăm A, B în pasul înapoi

Grafic ilustrând gradient checkpointing cu noduri și muchii

Antrenament eficient al modelelor AI cu PyTorch

Ce este gradient checkpointing?

  • Gradient checkpointing: reduce memoria selectând activările de salvat
  • Exemplu: calculăm A + B = C
    • Mai întâi calculăm A, B, apoi C
    • A, B nu sunt necesare în restul pasului înainte
  • Salvăm sau eliminăm A și B?
    • Fără gradient checkpointing: salvăm A, B
    • Gradient checkpointing: eliminăm A, B
    • Recalculăm A, B în pasul înapoi
    • Dacă B este costisitor de recalculat, îl salvăm

Grafic ilustrând gradient checkpointing cu noduri și muchii

Antrenament eficient al modelelor AI cu PyTorch

Trainer și Accelerator

Diagramă comparând ușurința de utilizare și posibilitatea de personalizare pentru Accelerator și Trainer.

Antrenament eficient al modelelor AI cu PyTorch

Trainer și Accelerator

Diagramă comparând ușurința de utilizare și posibilitatea de personalizare pentru Accelerator și Trainer.

Antrenament eficient al modelelor AI cu PyTorch

Gradient checkpointing cu Trainer

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







Antrenament eficient al modelelor AI cu PyTorch

Gradient checkpointing cu 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}
Antrenament eficient al modelelor AI cu PyTorch

De la Trainer la Accelerator

Diagramă comparând ușurința de utilizare și posibilitatea de personalizare pentru Accelerator și Trainer.

Antrenament eficient al modelelor AI cu PyTorch

Gradient checkpointing cu 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()
Antrenament eficient al modelelor AI cu PyTorch

Gradient checkpointing cu 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()
Antrenament eficient al modelelor AI cu PyTorch

SGD local îmbunătățește eficiența comunicării

 

 

Pictograme reprezentând eficiența memoriei, comunicării și calculului.

Antrenament eficient al modelelor AI cu PyTorch

Ce este SGD local?

Diagramă care arată cum funcționează SGD local prin sincronizarea gradienților după un număr de pași.

  • Fiecare dispozitiv calculează gradienții în paralel
Antrenament eficient al modelelor AI cu PyTorch

Ce este SGD local?

Diagramă care arată cum funcționează SGD local prin sincronizarea gradienților după un număr de pași.

  • Fiecare dispozitiv calculează gradienții în paralel
  • Sincronizarea gradienților: nodul coordonator actualizează parametrii modelului pe fiecare dispozitiv
  • SGD local: reduce frecvența sincronizării gradienților
Antrenament eficient al modelelor AI cu PyTorch

SGD local cu 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()

Antrenament eficient al modelelor AI cu PyTorch

SGD local cu 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()

Antrenament eficient al modelelor AI cu PyTorch

SGD local cu 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()
Antrenament eficient al modelelor AI cu PyTorch

Să exersăm!

Antrenament eficient al modelelor AI cu PyTorch

Preparing Video For Download...