Gradient checkpointing i lokalne SGD

Efektywne trenowanie modeli AI z PyTorch

Dennis Lee

Data Engineer, Amazon

Poprawa wydajności trenowania

 

 

Ikony reprezentujące wydajność pamięci, komunikacji i obliczeń.

Efektywne trenowanie modeli AI z PyTorch

Gradient checkpointing poprawia wydajność pamięci

 

 

Ikony reprezentujące wydajność pamięci, komunikacji i obliczeń.

Efektywne trenowanie modeli AI z PyTorch

Lokalne SGD poprawia wydajność komunikacji

 

 

Ikony reprezentujące wydajność pamięci, komunikacji i obliczeń.

Efektywne trenowanie modeli AI z PyTorch

Czym jest gradient checkpointing?

  • Gradient checkpointing: zmniejsza pamięć przez wybór aktywacji do zapisania
  • Przykład: oblicz A + B = C

Graf ilustrujący gradient checkpointing z węzłami i krawędziami

Efektywne trenowanie modeli AI z PyTorch

Czym jest gradient checkpointing?

  • Gradient checkpointing: zmniejsza pamięć przez wybór aktywacji do zapisania
  • Przykład: oblicz A + B = C
    • Najpierw oblicz A, B, potem C

Graf ilustrujący gradient checkpointing z węzłami i krawędziami

Efektywne trenowanie modeli AI z PyTorch

Czym jest gradient checkpointing?

  • Gradient checkpointing: zmniejsza pamięć przez wybór aktywacji do zapisania
  • Przykład: oblicz A + B = C
    • Najpierw oblicz A, B, potem C
    • A, B niepotrzebne w dalszym przebiegu w przód
  • Czy zapisać, czy usunąć A i B?

Graf ilustrujący gradient checkpointing z węzłami i krawędziami

Efektywne trenowanie modeli AI z PyTorch

Czym jest gradient checkpointing?

  • Gradient checkpointing: zmniejsza pamięć przez wybór aktywacji do zapisania
  • Przykład: oblicz A + B = C
    • Najpierw oblicz A, B, potem C
    • A, B niepotrzebne w dalszym przebiegu w przód
  • Czy zapisać, czy usunąć A i B?
    • Bez gradient checkpointing: zapisz A, B

Graf ilustrujący gradient checkpointing z węzłami i krawędziami

Efektywne trenowanie modeli AI z PyTorch

Czym jest gradient checkpointing?

  • Gradient checkpointing: zmniejsza pamięć przez wybór aktywacji do zapisania
  • Przykład: oblicz A + B = C
    • Najpierw oblicz A, B, potem C
    • A, B niepotrzebne w dalszym przebiegu w przód
  • Czy zapisać, czy usunąć A i B?
    • Bez gradient checkpointing: zapisz A, B
    • Gradient checkpointing: usuń A, B

Graf ilustrujący gradient checkpointing z węzłami i krawędziami

Efektywne trenowanie modeli AI z PyTorch

Czym jest gradient checkpointing?

  • Gradient checkpointing: zmniejsza pamięć przez wybór aktywacji do zapisania
  • Przykład: oblicz A + B = C
    • Najpierw oblicz A, B, potem C
    • A, B niepotrzebne w dalszym przebiegu w przód
  • Czy zapisać, czy usunąć A i B?
    • Bez gradient checkpointing: zapisz A, B
    • Gradient checkpointing: usuń A, B
    • Przelicz A, B podczas przebiegu wstecz

Graf ilustrujący gradient checkpointing z węzłami i krawędziami

Efektywne trenowanie modeli AI z PyTorch

Czym jest gradient checkpointing?

  • Gradient checkpointing: zmniejsza pamięć przez wybór aktywacji do zapisania
  • Przykład: oblicz A + B = C
    • Najpierw oblicz A, B, potem C
    • A, B niepotrzebne w dalszym przebiegu w przód
  • Czy zapisać, czy usunąć A i B?
    • Bez gradient checkpointing: zapisz A, B
    • Gradient checkpointing: usuń A, B
    • Przelicz A, B podczas przebiegu wstecz
    • Jeśli B jest kosztowne do przeliczenia, zapisz je

Graf ilustrujący gradient checkpointing z węzłami i krawędziami

Efektywne trenowanie modeli AI z PyTorch

Trainer i Accelerator

Wykres porównujący łatwość użycia i możliwości dostosowania dla Acceleratora i Trainera.

Efektywne trenowanie modeli AI z PyTorch

Trainer i Accelerator

Wykres porównujący łatwość użycia i możliwości dostosowania dla Acceleratora i Trainera.

Efektywne trenowanie modeli AI z PyTorch

Gradient checkpointing z Trainerem

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







Efektywne trenowanie modeli AI z PyTorch

Gradient checkpointing z Trainerem

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}
Efektywne trenowanie modeli AI z PyTorch

Od Trainera do Acceleratora

Wykres porównujący łatwość użycia i możliwości dostosowania dla Acceleratora i Trainera.

Efektywne trenowanie modeli AI z PyTorch

Gradient checkpointing z Acceleratorem

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()
Efektywne trenowanie modeli AI z PyTorch

Gradient checkpointing z Acceleratorem

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()
Efektywne trenowanie modeli AI z PyTorch

Lokalne SGD poprawia wydajność komunikacji

 

 

Ikony reprezentujące wydajność pamięci, komunikacji i obliczeń.

Efektywne trenowanie modeli AI z PyTorch

Czym jest lokalne SGD?

Diagram przedstawiający działanie lokalnego SGD przez synchronizację gradientów po określonej liczbie kroków.

  • Każde urządzenie oblicza gradienty równolegle
Efektywne trenowanie modeli AI z PyTorch

Czym jest lokalne SGD?

Diagram przedstawiający działanie lokalnego SGD przez synchronizację gradientów po określonej liczbie kroków.

  • Każde urządzenie oblicza gradienty równolegle
  • Synchronizacja gradientów: węzeł główny aktualizuje parametry modelu na każdym urządzeniu
  • Lokalne SGD: zmniejsza częstotliwość synchronizacji gradientów
Efektywne trenowanie modeli AI z PyTorch

Lokalne SGD z Acceleratorem





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()

Efektywne trenowanie modeli AI z PyTorch

Lokalne SGD z Acceleratorem

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()

Efektywne trenowanie modeli AI z PyTorch

Lokalne SGD z Acceleratorem

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()
Efektywne trenowanie modeli AI z PyTorch

Czas na ćwiczenia!

Efektywne trenowanie modeli AI z PyTorch

Preparing Video For Download...