Checkpointing gradient và Local SGD

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Dennis Lee

Data Engineer, Amazon

Nâng cao hiệu quả huấn luyện

 

 

Biểu tượng thể hiện hiệu quả bộ nhớ, hiệu quả truyền thông và hiệu quả tính toán.

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient cải thiện hiệu quả bộ nhớ

 

 

Biểu tượng thể hiện hiệu quả bộ nhớ, hiệu quả truyền thông và hiệu quả tính toán.

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Local SGD giải quyết hiệu quả truyền thông

 

 

Biểu tượng thể hiện hiệu quả bộ nhớ, hiệu quả truyền thông và hiệu quả tính toán.

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient là gì?

  • Checkpointing gradient: giảm bộ nhớ bằng cách chọn activation cần lưu
  • Ví dụ: tính A + B = C

Biểu đồ minh họa checkpointing gradient với các nút và cạnh

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient là gì?

  • Checkpointing gradient: giảm bộ nhớ bằng cách chọn activation cần lưu
  • Ví dụ: tính A + B = C
    • Tính A, B trước, rồi tính C

Biểu đồ minh họa checkpointing gradient với các nút và cạnh

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient là gì?

  • Checkpointing gradient: giảm bộ nhớ bằng cách chọn activation cần lưu
  • Ví dụ: tính A + B = C
    • Tính A, B trước, rồi tính C
    • A, B không cần cho phần còn lại của forward pass
  • Ta nên lưu hay xóa A và B?

Biểu đồ minh họa checkpointing gradient với các nút và cạnh

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient là gì?

  • Checkpointing gradient: giảm bộ nhớ bằng cách chọn activation cần lưu
  • Ví dụ: tính A + B = C
    • Tính A, B trước, rồi tính C
    • A, B không cần cho phần còn lại của forward pass
  • Ta nên lưu hay xóa A và B?
    • Không dùng checkpointing: lưu A, B

Biểu đồ minh họa checkpointing gradient với các nút và cạnh

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient là gì?

  • Checkpointing gradient: giảm bộ nhớ bằng cách chọn activation cần lưu
  • Ví dụ: tính A + B = C
    • Tính A, B trước, rồi tính C
    • A, B không cần cho phần còn lại của forward pass
  • Ta nên lưu hay xóa A và B?
    • Không dùng checkpointing: lưu A, B
    • Dùng checkpointing: xóa A, B

Biểu đồ minh họa checkpointing gradient với các nút và cạnh

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient là gì?

  • Checkpointing gradient: giảm bộ nhớ bằng cách chọn activation cần lưu
  • Ví dụ: tính A + B = C
    • Tính A, B trước, rồi tính C
    • A, B không cần cho phần còn lại của forward pass
  • Ta nên lưu hay xóa A và B?
    • Không dùng checkpointing: lưu A, B
    • Dùng checkpointing: xóa A, B
    • Tính lại A, B trong backward pass

Biểu đồ minh họa checkpointing gradient với các nút và cạnh

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient là gì?

  • Checkpointing gradient: giảm bộ nhớ bằng cách chọn activation cần lưu
  • Ví dụ: tính A + B = C
    • Tính A, B trước, rồi tính C
    • A, B không cần cho phần còn lại của forward pass
  • Ta nên lưu hay xóa A và B?
    • Không dùng checkpointing: lưu A, B
    • Dùng checkpointing: xóa A, B
    • Tính lại A, B trong backward pass
    • Nếu B tốn kém khi tính lại, hãy lưu B

Biểu đồ minh họa checkpointing gradient với các nút và cạnh

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Trainer và Accelerator

Biểu đồ so sánh dễ dùng và khả năng tùy chỉnh giữa Accelerator và Trainer.

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Trainer và Accelerator

Biểu đồ so sánh dễ dùng và khả năng tùy chỉnh giữa Accelerator và Trainer.

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient với Trainer

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







Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient với 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}
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Từ Trainer sang Accelerator

Biểu đồ so sánh dễ dùng và khả năng tùy chỉnh giữa Accelerator và Trainer.

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient với 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()
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Checkpointing gradient với 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()
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Local SGD cải thiện hiệu quả truyền thông

 

 

Biểu tượng thể hiện hiệu quả bộ nhớ, hiệu quả truyền thông và hiệu quả tính toán.

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Local SGD là gì?

Sơ đồ cho thấy Local SGD đồng bộ gradient sau một số bước nhất định.

  • Mỗi thiết bị tính gradient song song
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Local SGD là gì?

Sơ đồ cho thấy Local SGD đồng bộ gradient sau một số bước nhất định.

  • Mỗi thiết bị tính gradient song song
  • Đồng bộ gradient: nút điều phối cập nhật tham số mô hình trên mỗi thiết bị
  • Local SGD: giảm tần suất đồng bộ gradient
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Local SGD với Accelerator





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

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Local SGD với Accelerator

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

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Local SGD với Accelerator

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()
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Ayo berlatih!

Huấn luyện Mô hình AI Hiệu quả với PyTorch

Preparing Video For Download...