ПочатиПочніть безкоштовно

Local SGD з Accelerator

Ви вже реалізували акумуляцію градієнтів і gradient checkpointing, щоб оптимізувати використання пам’яті у своїй моделі перекладу. Навчання все ще трохи повільне, тож ви вирішили додати local SGD до циклу навчання, щоб підвищити ефективність обміну даними між пристроями. Побудуйте цикл навчання з local SGD!

model, train_dataloader і accelerator уже визначені, а LocalSGD імпортовано.

Ця вправа є частиною курсу

Ефективне тренування моделей ШІ з PyTorch

Переглянути курс

Інструкції до вправи

  • Встановіть local_sgd_steps, щоб синхронізувати градієнти кожні вісім кроків.
  • Виконайте крок менеджера контексту local SGD.

Інтерактивна практична вправа

Спробуйте виконати цю вправу, доповнивши цей зразок коду.

# 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.____()
Редагувати та запускати код