ÎncepețiÎncepe gratuit

Local SGD cu Accelerator

Ai implementat acumularea gradienților și checkpoint-ul gradienților pentru a optimiza utilizarea memoriei în modelul tău de traducere automată. Antrenarea este încă puțin lentă, așa că decizi să adaugi local SGD în bucla de antrenare pentru a îmbunătăți eficiența comunicării între dispozitive. Construiește bucla de antrenare cu local SGD!

model, train_dataloader și accelerator sunt predefinite, iar LocalSGD a fost importat.

Acest exercițiu face parte din cursul

Antrenament eficient al modelelor AI cu PyTorch

Vezi cursul

Instrucțiuni pentru exercițiu

  • Setează local_sgd_steps pentru a sincroniza gradienții la fiecare opt pași.
  • Apelează metoda step a managerului de context local SGD.

Exercițiu interactiv practic

Încearcă acest exercițiu completând acest cod de exemplu.

# Set up a context manager to synchronize gradients every eight steps
with LocalSGD(accelerator=accelerator, model=model, local_sgd_steps=____, enabled=True) as local_sgd:
    for batch in train_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()
            # Step the local SGD context manager
            local_sgd.____()
Editează și rulează codul