Huấn luyện cân bằng với AdamW

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

Dennis Lee

Data Engineer, Amazon

Huấn luyện hiệu quả

 

 

Sơ đồ hiển thị các chủ đề chương trong khóa học, tô đậm mục bộ tối ưu.

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

Bộ tối ưu cho hiệu quả huấn luyện

 

 

Sơ đồ mô tả ba bộ tối ưu: AdamW, Adafactor và Adam 8-bit.

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

Bộ tối ưu cho hiệu quả huấn luyện

 

 

Sơ đồ mô tả ba bộ tối ưu: AdamW, Adafactor và Adam 8-bit.

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

Bộ tối ưu cho hiệu quả huấn luyện

 

 

Sơ đồ mô tả ba bộ tối ưu: AdamW, Adafactor và Adam 8-bit.

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

Đánh đổi của bộ tối ưu

Sơ đồ cho thấy đánh đổi giữa số tham số và độ chính xác của AdamW, Adafactor và Adam 8-bit.

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

Đánh đổi của bộ tối ưu

Sơ đồ cho thấy đánh đổi giữa số tham số và độ chính xác của AdamW, Adafactor và Adam 8-bit.

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

Đánh đổi của bộ tối ưu

Sơ đồ cho thấy đánh đổi giữa số tham số và độ chính xác của AdamW, Adafactor và Adam 8-bit.

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

Đánh đổi của bộ tối ưu

Sơ đồ cho thấy đánh đổi giữa số tham số và độ chính xác của AdamW, Adafactor và Adam 8-bit.

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

AdamW hoạt động thế nào?

Sơ đồ minh họa cách AdamW hoạt động.

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

AdamW hoạt động thế nào?

Sơ đồ minh họa cách AdamW hoạt động.

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

AdamW hoạt động thế nào?

Sơ đồ minh họa cách AdamW hoạt động.

  • Tính trung bình trượt mũ (EMA) của gradient
Huấn luyện Mô hình AI Hiệu quả với PyTorch

AdamW hoạt động thế nào?

Sơ đồ minh họa cách AdamW hoạt động.

  • Tính trung bình trượt mũ (EMA) của gradient
  • Tính EMA của bình phương gradient
Huấn luyện Mô hình AI Hiệu quả với PyTorch

AdamW hoạt động thế nào?

Sơ đồ minh họa cách AdamW hoạt động.

  • Tính trung bình trượt mũ (EMA) của gradient
  • Tính EMA của bình phương gradient
Huấn luyện Mô hình AI Hiệu quả với PyTorch

AdamW hoạt động thế nào?

Sơ đồ minh họa cách AdamW hoạt động.

  • Tính trung bình trượt mũ (EMA) của gradient
  • Tính EMA của bình phương gradient
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Bộ nhớ dùng bởi AdamW

Sơ đồ hiển thị các đại lượng trong tính toán AdamW: gradient tham số, EMA của gradient và EMA của bình phương gradient.

  • Mỗi ô là một tham số, mỗi màu là một trạng thái
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Bộ nhớ dùng bởi AdamW

Sơ đồ hiển thị các đại lượng trong tính toán AdamW: gradient tham số, EMA của gradient và EMA của bình phương gradient.

  • Mỗi ô là một tham số, mỗi màu là một trạng thái
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Bộ nhớ dùng bởi AdamW

Sơ đồ hiển thị các đại lượng trong tính toán AdamW: gradient tham số, EMA của gradient và EMA của bình phương gradient.

  • Mỗi ô là một tham số, mỗi màu là một trạng thái
  • Bộ nhớ mỗi tham số = 8 byte = 4 byte mỗi trạng thái * 2 trạng thái
  • Tổng bộ nhớ = Bộ nhớ mỗi tham số (8 byte) * Số tham số
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Ước tính bộ nhớ dùng bởi AdamW

model = AutoModelForSequenceClassification.from_pretrained(
    "distilbert-base-cased", return_dict=True)

num_parameters = sum(p.numel() for p in model.parameters()) print(f"Number of model parameters: {num_parameters:,}")
Number of model parameters: 65,783,042
estimated_memory = num_parameters * 8 / (1024 ** 2)
print(f"Estimated memory usage of AdamW: {estimated_memory:.0f} MB")
Estimated memory usage of AdamW: 502 MB
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Trainer và Accelerator

Sơ đồ cho thấy đánh đổi giữa khả năng tùy biến và dễ dùng của Accelerator và Trainer.

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

Triển khai AdamW với Trainer

from torch.optim import AdamW

optimizer = AdamW(params=model.parameters())


trainer = Trainer(model=model, args=training_args, train_dataset=train_dataset, eval_dataset=validation_dataset, compute_metrics=compute_metrics, optimizers=(optimizer, lr_scheduler))
trainer.train()
{'epoch': 1.0, 'eval_accuracy': 0.7, 'eval_f1': 0.8}
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Triển khai AdamW với Accelerator

from torch.optim import AdamW

optimizer = AdamW(params=model.parameters())


for batch in train_dataloader: 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() print(f"Loss = {loss}")
Loss = 0.7
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Kiểm tra trạng thái bộ tối ưu

optimizer_state = optimizer.state.values()
print(optimizer_state)
dict_values([{'step': tensor(3.),

'exp_avg': tensor([[0., 0., 0., ..., 0., 0., 0.], ...]),
'exp_avg_sq': tensor([[0., 0., 0., ..., 0., 0., 0.], ...])}, ...])
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Tính kích thước bộ tối ưu

def compute_optimizer_size(optimizer_state):
    total_size_megabytes, total_num_elements = 0, 0

for params in optimizer_state:
for name, tensor in params.items(): tensor = torch.tensor(tensor)
num_elements = tensor.numel()
element_size = tensor.element_size()
total_num_elements += num_elements
total_size_megabytes += num_elements * element_size / (1024 ** 2)
return total_size_megabytes, total_num_elements
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Tính kích thước bộ tối ưu

total_size_megabytes, total_num_elements = \
    compute_optimizer_size(trainer.optimizer.state.values())
print(f"Number of optimizer parameters: {total_num_elements:,}")
Number of optimizer parameters: 131,566,188
print(f"Optimizer size: {total_size_megabytes:.0f} MB")
Optimizer size: 502 MB
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...