Gradientkontrollpunkter och lokal SGD

Effektiv AI-modellträning med PyTorch

Dennis Lee

Data Engineer, Amazon

Förbättra träningseffektivitet

 

 

Ikoner som representerar minneseffektivitet, kommunikationseffektivitet och beräkningseffektivitet.

Effektiv AI-modellträning med PyTorch

Gradientkontrollpunkter förbättrar minneseffektivitet

 

 

Ikoner som representerar minneseffektivitet, kommunikationseffektivitet och beräkningseffektivitet.

Effektiv AI-modellträning med PyTorch

Lokal SGD adresserar kommunikationseffektivitet

 

 

Ikoner som representerar minneseffektivitet, kommunikationseffektivitet och beräkningseffektivitet.

Effektiv AI-modellträning med PyTorch

Vad är gradientkontrollpunkter?

  • Gradientkontrollpunkt: minska minne genom att välja vilka aktiveringar som sparas
  • Exempel: beräkna A + B = C

Graf som illustrerar gradientkontrollpunkter med noder och kanter

Effektiv AI-modellträning med PyTorch

Vad är gradientkontrollpunkter?

  • Gradientkontrollpunkt: minska minne genom att välja vilka aktiveringar som sparas
  • Exempel: beräkna A + B = C
    • Beräkna först A, B, sedan C

Graf som illustrerar gradientkontrollpunkter med noder och kanter

Effektiv AI-modellträning med PyTorch

Vad är gradientkontrollpunkter?

  • Gradientkontrollpunkt: minska minne genom att välja vilka aktiveringar som sparas
  • Exempel: beräkna A + B = C
    • Beräkna först A, B, sedan C
    • A, B behövs inte för resten av framåtpasset
  • Ska vi spara eller ta bort A och B?

Graf som illustrerar gradientkontrollpunkter med noder och kanter

Effektiv AI-modellträning med PyTorch

Vad är gradientkontrollpunkter?

  • Gradientkontrollpunkt: minska minne genom att välja vilka aktiveringar som sparas
  • Exempel: beräkna A + B = C
    • Beräkna först A, B, sedan C
    • A, B behövs inte för resten av framåtpasset
  • Ska vi spara eller ta bort A och B?
    • Utan gradientkontrollpunkt: spara A, B

Graf som illustrerar gradientkontrollpunkter med noder och kanter

Effektiv AI-modellträning med PyTorch

Vad är gradientkontrollpunkter?

  • Gradientkontrollpunkt: minska minne genom att välja vilka aktiveringar som sparas
  • Exempel: beräkna A + B = C
    • Beräkna först A, B, sedan C
    • A, B behövs inte för resten av framåtpasset
  • Ska vi spara eller ta bort A och B?
    • Utan gradientkontrollpunkt: spara A, B
    • Gradientkontrollpunkt: ta bort A, B

Graf som illustrerar gradientkontrollpunkter med noder och kanter

Effektiv AI-modellträning med PyTorch

Vad är gradientkontrollpunkter?

  • Gradientkontrollpunkt: minska minne genom att välja vilka aktiveringar som sparas
  • Exempel: beräkna A + B = C
    • Beräkna först A, B, sedan C
    • A, B behövs inte för resten av framåtpasset
  • Ska vi spara eller ta bort A och B?
    • Utan gradientkontrollpunkt: spara A, B
    • Gradientkontrollpunkt: ta bort A, B
    • Beräkna om A, B under bakåtpasset

Graf som illustrerar gradientkontrollpunkter med noder och kanter

Effektiv AI-modellträning med PyTorch

Vad är gradientkontrollpunkter?

  • Gradientkontrollpunkt: minska minne genom att välja vilka aktiveringar som sparas
  • Exempel: beräkna A + B = C
    • Beräkna först A, B, sedan C
    • A, B behövs inte för resten av framåtpasset
  • Ska vi spara eller ta bort A och B?
    • Utan gradientkontrollpunkt: spara A, B
    • Gradientkontrollpunkt: ta bort A, B
    • Beräkna om A, B under bakåtpasset
    • Om B är kostsamt att beräkna om, spara det

Graf som illustrerar gradientkontrollpunkter med noder och kanter

Effektiv AI-modellträning med PyTorch

Trainer och Accelerator

Diagram som jämför användarvänlighet och anpassningsbarhet för Accelerator och Trainer.

Effektiv AI-modellträning med PyTorch

Trainer och Accelerator

Diagram som jämför användarvänlighet och anpassningsbarhet för Accelerator och Trainer.

Effektiv AI-modellträning med PyTorch

Gradientkontrollpunkter med Trainer

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







Effektiv AI-modellträning med PyTorch

Gradientkontrollpunkter med 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}
Effektiv AI-modellträning med PyTorch

Från Trainer till Accelerator

Diagram som jämför användarvänlighet och anpassningsbarhet för Accelerator och Trainer.

Effektiv AI-modellträning med PyTorch

Gradientkontrollpunkter med 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()
Effektiv AI-modellträning med PyTorch

Gradientkontrollpunkter med 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()
Effektiv AI-modellträning med PyTorch

Lokal SGD förbättrar kommunikationseffektivitet

 

 

Ikoner som representerar minneseffektivitet, kommunikationseffektivitet och beräkningseffektivitet.

Effektiv AI-modellträning med PyTorch

Vad är lokal SGD?

Diagram som visar hur lokal SGD fungerar genom att synkronisera gradienter efter ett visst antal steg.

  • Varje enhet beräknar gradienter parallellt
Effektiv AI-modellträning med PyTorch

Vad är lokal SGD?

Diagram som visar hur lokal SGD fungerar genom att synkronisera gradienter efter ett visst antal steg.

  • Varje enhet beräknar gradienter parallellt
  • Gradientsynkronisering: drivernoden uppdaterar modellparametrar på varje enhet
  • Lokal SGD: minska frekvensen av gradientsynkronisering
Effektiv AI-modellträning med PyTorch

Lokal SGD med 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()

Effektiv AI-modellträning med PyTorch

Lokal SGD med 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()

Effektiv AI-modellträning med PyTorch

Lokal SGD med 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()
Effektiv AI-modellträning med PyTorch

Låt oss öva!

Effektiv AI-modellträning med PyTorch

Preparing Video For Download...