Локальный SGD с Accelerator
Вы уже реализовали накопление градиентов и контрольные точки градиентов, чтобы оптимизировать использование памяти в модели перевода текста. Обучение по-прежнему идёт медленновато, поэтому вы решаете добавить локальный SGD в цикл обучения — это позволит повысить эффективность обмена данными между устройствами. Создайте цикл обучения с локальным SGD!
model, train_dataloader и accelerator уже определены, а LocalSGD импортирован.
Это упражнение является частью курса
Эффективное обучение моделей ИИ с PyTorch
Инструкции к упражнению
- Установите
local_sgd_stepsтак, чтобы синхронизация градиентов выполнялась каждые восемь шагов. - Выполните шаг менеджера контекста локального 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.____()