Huấn luyện độ chính xác hỗn hợp với Adam 8-bit

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

Dennis Lee

Data Engineer, Amazon

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

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

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

Adam 8-bit hoạt động thế nào?

Sơ đồ mô tả các bước của Adam 8-bit.

  • Lưu tham số ở FP8; tối ưu ở FP32
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Adam 8-bit hoạt động thế nào?

Sơ đồ mô tả các bước của Adam 8-bit.

  • Lưu tham số ở FP8; tối ưu ở FP32
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Adam 8-bit hoạt động thế nào?

Sơ đồ mô tả các bước của Adam 8-bit.

  • Lưu tham số ở FP8; tối ưu ở FP32
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Adam 8-bit hoạt động thế nào?

Sơ đồ mô tả các bước của Adam 8-bit.

  • Lưu tham số ở FP8; tối ưu ở FP32
  • EMA: trung bình trượt mũ
  • Tính EMA của gradient và gradient bình phương
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Adam 8-bit hoạt động thế nào?

Sơ đồ mô tả các bước của Adam 8-bit.

  • Lưu tham số ở FP8; tối ưu ở FP32
  • EMA: trung bình trượt mũ
  • Tính EMA của gradient và gradient bình phương
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Adam 8-bit tiết kiệm bộ nhớ thế nào?

Sơ đồ thể hiện gradient tham số, EMA của gradient, và EMA của gradient bình phương.

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

Adam 8-bit tiết kiệm bộ nhớ thế nào?

Sơ đồ thể hiện gradient tham số, EMA của gradient, và EMA của gradient bình phương.

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

Adam 8-bit tiết kiệm bộ nhớ thế nào?

Sơ đồ thể hiện gradient tham số, EMA của gradient, và EMA của gradient bình phương.

  • Mỗi ô là một tham số, mỗi màu là một trạng thái
  • Bộ nhớ mỗi tham số = 2 byte = 1 byte/trạng thái * 2 trạng thái
  • Tổng bộ nhớ = Bộ nhớ mỗi tham số (2 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 cho Adam 8-bit

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 * 2 / (1024 ** 2)
print(f"Estimated memory usage of 8-bit Adam: {estimated_memory:.0f} MB")
Estimated memory usage of 8-bit Adam: 125 MB
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Thiết lập bộ tối ưu Adam 8-bit

import bitsandbytes as bnb
from torch import nn
from transformers.trainer_pt_utils import get_parameter_names


args = TrainingArguments(output_dir="./results")
decay_parameters = get_parameter_names(model, [nn.LayerNorm])
decay_parameters = [name for name in decay_parameters if "bias" not in name]
  • Weight decay giúp tránh overfitting
  • decay_parameters: chỉ định tham số áp dụng weight decay
  • get_parameter_names: trả về tên tham số; bỏ qua các lớp nn.LayerNorm
  • Loại tham số bias khỏi decay_parameters
  • Không áp dụng weight decay cho lớp chuẩn hóa và bias
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Thiết lập bộ tối ưu Adam 8-bit

optimizer_grouped_parameters = [{"params": [p for n, p in model.named_parameters() 
                                            if n in decay_parameters],
                                 "weight_decay": args.weight_decay},

{"params": [p for n, p in model.named_parameters() if n not in decay_parameters], "weight_decay": 0.0}]
adam_bnb_optim = bnb.optim.Adam8bit(optimizer_grouped_parameters,
betas=(args.adam_beta1, args.adam_beta2),
eps=args.adam_epsilon,
lr=args.learning_rate)
  • optimizer_grouped_parameters: Một nhóm áp dụng weight decay; nhóm còn lại không
  • beta1, beta2: Hệ số suy giảm moment 1 và 2; cao hơn = ổn định hơn, học chậm hơn
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Triển khai Adam 8-bit với Trainer

trainer = Trainer(model=model,
                  args=training_args,
                  train_dataset=train_dataset,
                  eval_dataset=validation_dataset,
                  optimizers=(adam_bnb_optim, None),
                  compute_metrics=compute_metrics)

trainer.train()
{'epoch': 1.0, 'eval_loss': 0.63, 'eval_accuracy': 0.67, 'eval_f1': 0.62}
{'epoch': 2.0, 'eval_loss': 0.61, 'eval_accuracy': 0.71, 'eval_f1': 0.66}
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Triển khai Adam 8-bit với Accelerator

model, adam_bnb_optim, train_dataloader, lr_scheduler = \
    accelerator.prepare(model, adam_bnb_optim, train_dataloader, lr_scheduler)


for batch in train_dataloader: inputs, targets = batch["input_ids"], batch["labels"] outputs = model(inputs, labels=targets) loss = outputs.loss accelerator.backward(loss) adam_bnb_optim.step() lr_scheduler.step() adam_bnb_optim.zero_grad() print(f"Loss = {loss}")
Loss = 0.75
Huấn luyện Mô hình AI Hiệu quả với PyTorch

Tính bộ nhớ dùng cho Adam 8-bit

total_size_megabytes, total_num_elements = \
    compute_optimizer_size(trainer.optimizer.state.values())
print(f"Number of 8-bit Adam parameters: {total_num_elements:,}")
print(f"8-bit Adam size: {total_size_megabytes:.0f} MB")
Number of 8-bit Adam parameters: 131,566,188
8-bit Adam size: 128 MB
  • So với AdamW: Adam 8-bit dùng 1/4 bộ nhớ
Number of AdamW parameters: 131,566,188
AdamW size: 502 MB
Huấn luyện Mô hình AI Hiệu quả với PyTorch

¡Vamos a practicar!

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

Preparing Video For Download...