Gradient checkpointing e SGD local

Treinamento Eficiente de Modelos de IA com PyTorch

Dennis Lee

Data Engineer, Amazon

Melhorando a eficiência do treino

 

 

Ícones representando eficiência de memória, de comunicação e computacional.

Treinamento Eficiente de Modelos de IA com PyTorch

Gradient checkpointing melhora a eficiência de memória

 

 

Ícones representando eficiência de memória, de comunicação e computacional.

Treinamento Eficiente de Modelos de IA com PyTorch

SGD local melhora a eficiência de comunicação

 

 

Ícones representando eficiência de memória, de comunicação e computacional.

Treinamento Eficiente de Modelos de IA com PyTorch

O que é gradient checkpointing?

  • Gradient checkpointing: reduz memória escolhendo quais ativações salvar
  • Exemplo: calcule A + B = C

Gráfico ilustrando gradient checkpointing com nós e arestas

Treinamento Eficiente de Modelos de IA com PyTorch

O que é gradient checkpointing?

  • Gradient checkpointing: reduz memória escolhendo quais ativações salvar
  • Exemplo: calcule A + B = C
    • Primeiro calcule A e B, depois calcule C

Gráfico ilustrando gradient checkpointing com nós e arestas

Treinamento Eficiente de Modelos de IA com PyTorch

O que é gradient checkpointing?

  • Gradient checkpointing: reduz memória escolhendo quais ativações salvar
  • Exemplo: calcule A + B = C
    • Primeiro calcule A e B, depois calcule C
    • A e B não são usados no resto do forward
  • Devemos salvar ou descartar A e B?

Gráfico ilustrando gradient checkpointing com nós e arestas

Treinamento Eficiente de Modelos de IA com PyTorch

O que é gradient checkpointing?

  • Gradient checkpointing: reduz memória escolhendo quais ativações salvar
  • Exemplo: calcule A + B = C
    • Primeiro calcule A e B, depois calcule C
    • A e B não são usados no resto do forward
  • Devemos salvar ou descartar A e B?
    • Sem gradient checkpointing: salve A e B

Gráfico ilustrando gradient checkpointing com nós e arestas

Treinamento Eficiente de Modelos de IA com PyTorch

O que é gradient checkpointing?

  • Gradient checkpointing: reduz memória escolhendo quais ativações salvar
  • Exemplo: calcule A + B = C
    • Primeiro calcule A e B, depois calcule C
    • A e B não são usados no resto do forward
  • Devemos salvar ou descartar A e B?
    • Sem gradient checkpointing: salve A e B
    • Com gradient checkpointing: descarte A e B

Gráfico ilustrando gradient checkpointing com nós e arestas

Treinamento Eficiente de Modelos de IA com PyTorch

O que é gradient checkpointing?

  • Gradient checkpointing: reduz memória escolhendo quais ativações salvar
  • Exemplo: calcule A + B = C
    • Primeiro calcule A e B, depois calcule C
    • A e B não são usados no resto do forward
  • Devemos salvar ou descartar A e B?
    • Sem gradient checkpointing: salve A e B
    • Com gradient checkpointing: descarte A e B
    • Recalcule A e B no backward

Gráfico ilustrando gradient checkpointing com nós e arestas

Treinamento Eficiente de Modelos de IA com PyTorch

O que é gradient checkpointing?

  • Gradient checkpointing: reduz memória escolhendo quais ativações salvar
  • Exemplo: calcule A + B = C
    • Primeiro calcule A e B, depois calcule C
    • A e B não são usados no resto do forward
  • Devemos salvar ou descartar A e B?
    • Sem gradient checkpointing: salve A e B
    • Com gradient checkpointing: descarte A e B
    • Recalcule A e B no backward
    • Se B for caro de recalcular, salve-o

Gráfico ilustrando gradient checkpointing com nós e arestas

Treinamento Eficiente de Modelos de IA com PyTorch

Trainer e Accelerator

Gráfico comparando facilidade de uso vs. possibilidade de customização para Accelerator e Trainer.

Treinamento Eficiente de Modelos de IA com PyTorch

Trainer e Accelerator

Gráfico comparando facilidade de uso vs. possibilidade de customização para Accelerator e Trainer.

Treinamento Eficiente de Modelos de IA com PyTorch

Gradient checkpointing com Trainer

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







Treinamento Eficiente de Modelos de IA com PyTorch

Gradient checkpointing com 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}
Treinamento Eficiente de Modelos de IA com PyTorch

Do Trainer ao Accelerator

Gráfico comparando facilidade de uso vs. possibilidade de customização para Accelerator e Trainer.

Treinamento Eficiente de Modelos de IA com PyTorch

Gradient checkpointing com 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()
Treinamento Eficiente de Modelos de IA com PyTorch

Gradient checkpointing com 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()
Treinamento Eficiente de Modelos de IA com PyTorch

SGD local melhora a eficiência de comunicação

 

 

Ícones representando eficiência de memória, de comunicação e computacional.

Treinamento Eficiente de Modelos de IA com PyTorch

O que é SGD local?

Diagrama mostrando como o SGD local sincroniza gradientes após certo número de passos.

  • Cada dispositivo calcula gradientes em paralelo
Treinamento Eficiente de Modelos de IA com PyTorch

O que é SGD local?

Diagrama mostrando como o SGD local sincroniza gradientes após certo número de passos.

  • Cada dispositivo calcula gradientes em paralelo
  • Sincronização de gradientes: o nó driver atualiza os parâmetros do modelo em cada dispositivo
  • SGD local: reduza a frequência de sincronização de gradientes
Treinamento Eficiente de Modelos de IA com PyTorch

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

Treinamento Eficiente de Modelos de IA com PyTorch

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

Treinamento Eficiente de Modelos de IA com PyTorch

SGD local com 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()
Treinamento Eficiente de Modelos de IA com PyTorch

Vamos praticar!

Treinamento Eficiente de Modelos de IA com PyTorch

Preparing Video For Download...