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
Instrucțiuni pentru exercițiu
- Setează
local_sgd_stepspentru 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.____()