梯度檢查點與本地 SGD

使用 PyTorch 高效訓練 AI 模型

Dennis Lee

Data Engineer, Amazon

提升訓練效率

 

 

代表記憶體效率、通訊效率與運算效率的圖示。

使用 PyTorch 高效訓練 AI 模型

梯度檢查點提升記憶體效率

 

 

代表記憶體效率、通訊效率與運算效率的圖示。

使用 PyTorch 高效訓練 AI 模型

本地 SGD 改善通訊效率

 

 

代表記憶體效率、通訊效率與運算效率的圖示。

使用 PyTorch 高效訓練 AI 模型

什麼是梯度檢查點?

  • 梯度檢查點:選擇要保留的啟動值以降低記憶體
  • 範例:計算 A + B = C

以節點與邊說明梯度檢查點的圖示

使用 PyTorch 高效訓練 AI 模型

什麼是梯度檢查點?

  • 梯度檢查點:選擇要保留的啟動值以降低記憶體
  • 範例:計算 A + B = C
    • 先算出 A、B,再算 C

以節點與邊說明梯度檢查點的圖示

使用 PyTorch 高效訓練 AI 模型

什麼是梯度檢查點?

  • 梯度檢查點:選擇要保留的啟動值以降低記憶體
  • 範例:計算 A + B = C
    • 先算出 A、B,再算 C
    • 之後的前向傳播不再需要 A、B
  • 我們該保留還是移除 A、B?

以節點與邊說明梯度檢查點的圖示

使用 PyTorch 高效訓練 AI 模型

什麼是梯度檢查點?

  • 梯度檢查點:選擇要保留的啟動值以降低記憶體
  • 範例:計算 A + B = C
    • 先算出 A、B,再算 C
    • 之後的前向傳播不再需要 A、B
  • 我們該保留還是移除 A、B?
    • 未用梯度檢查點:保留 A、B

以節點與邊說明梯度檢查點的圖示

使用 PyTorch 高效訓練 AI 模型

什麼是梯度檢查點?

  • 梯度檢查點:選擇要保留的啟動值以降低記憶體
  • 範例:計算 A + B = C
    • 先算出 A、B,再算 C
    • 之後的前向傳播不再需要 A、B
  • 我們該保留還是移除 A、B?
    • 未用梯度檢查點:保留 A、B
    • 使用梯度檢查點:移除 A、B

以節點與邊說明梯度檢查點的圖示

使用 PyTorch 高效訓練 AI 模型

什麼是梯度檢查點?

  • 梯度檢查點:選擇要保留的啟動值以降低記憶體
  • 範例:計算 A + B = C
    • 先算出 A、B,再算 C
    • 之後的前向傳播不再需要 A、B
  • 我們該保留還是移除 A、B?
    • 未用梯度檢查點:保留 A、B
    • 使用梯度檢查點:移除 A、B
    • 在反向傳播時重新計算 A、B

以節點與邊說明梯度檢查點的圖示

使用 PyTorch 高效訓練 AI 模型

什麼是梯度檢查點?

  • 梯度檢查點:選擇要保留的啟動值以降低記憶體
  • 範例:計算 A + B = C
    • 先算出 A、B,再算 C
    • 之後的前向傳播不再需要 A、B
  • 我們該保留還是移除 A、B?
    • 未用梯度檢查點:保留 A、B
    • 使用梯度檢查點:移除 A、B
    • 在反向傳播時重新計算 A、B
    • 若 B 重算成本高,則保留 B

以節點與邊說明梯度檢查點的圖示

使用 PyTorch 高效訓練 AI 模型

Trainer 與 Accelerator

比較 Accelerator 與 Trainer 在易用性與自訂能力的圖表。

使用 PyTorch 高效訓練 AI 模型

Trainer 與 Accelerator

比較 Accelerator 與 Trainer 在易用性與自訂能力的圖表。

使用 PyTorch 高效訓練 AI 模型

在 Trainer 中使用梯度檢查點

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







使用 PyTorch 高效訓練 AI 模型

在 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}
使用 PyTorch 高效訓練 AI 模型

從 Trainer 轉到 Accelerator

比較 Accelerator 與 Trainer 在易用性與自訂能力的圖表。

使用 PyTorch 高效訓練 AI 模型

在 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()
使用 PyTorch 高效訓練 AI 模型

在 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()
使用 PyTorch 高效訓練 AI 模型

本地 SGD 提升通訊效率

 

 

代表記憶體效率、通訊效率與運算效率的圖示。

使用 PyTorch 高效訓練 AI 模型

什麼是本地 SGD?

示意圖:本地 SGD 透過每隔數步同步梯度來運作。

  • 每個裝置並行計算梯度
使用 PyTorch 高效訓練 AI 模型

什麼是本地 SGD?

示意圖:本地 SGD 透過每隔數步同步梯度來運作。

  • 每個裝置並行計算梯度
  • 梯度同步:Driver 節點更新各裝置的模型參數
  • 本地 SGD:降低梯度同步頻率
使用 PyTorch 高效訓練 AI 模型

在 Accelerator 中使用本地 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()

使用 PyTorch 高效訓練 AI 模型

在 Accelerator 中使用本地 SGD

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()

使用 PyTorch 高效訓練 AI 模型

在 Accelerator 中使用本地 SGD

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()
使用 PyTorch 高效訓練 AI 模型

一起來練習吧!

使用 PyTorch 高效訓練 AI 模型

Preparing Video For Download...