การฝึกแบบ Mixed Precision ด้วย 8-bit Adam

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Dennis Lee

Data Engineer, Amazon

Optimizer เพื่อประสิทธิภาพการฝึก

แผนภาพแสดงการเปรียบเทียบระหว่างจำนวนพารามิเตอร์และความแม่นยำสำหรับ AdamW, Adafactor และ 8-bit Adam

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Optimizer เพื่อประสิทธิภาพการฝึก

แผนภาพแสดงการเปรียบเทียบระหว่างจำนวนพารามิเตอร์และความแม่นยำสำหรับ AdamW, Adafactor และ 8-bit Adam

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

8-bit Adam ทำงานอย่างไร?

แผนภาพแสดงขั้นตอนของ 8-bit Adam

  • จัดเก็บพารามิเตอร์ใน FP8; ปรับปรุงค่าใน FP32
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

8-bit Adam ทำงานอย่างไร?

แผนภาพแสดงขั้นตอนของ 8-bit Adam

  • จัดเก็บพารามิเตอร์ใน FP8; ปรับปรุงค่าใน FP32
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

8-bit Adam ทำงานอย่างไร?

แผนภาพแสดงขั้นตอนของ 8-bit Adam

  • จัดเก็บพารามิเตอร์ใน FP8; ปรับปรุงค่าใน FP32
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

8-bit Adam ทำงานอย่างไร?

แผนภาพแสดงขั้นตอนของ 8-bit Adam

  • จัดเก็บพารามิเตอร์ใน FP8; ปรับปรุงค่าใน FP32
  • EMA: ค่าเฉลี่ยเคลื่อนที่แบบเอกซ์โพเนนเชียล
  • คำนวณ EMA ของ gradient และ gradient กำลังสอง
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

8-bit Adam ทำงานอย่างไร?

แผนภาพแสดงขั้นตอนของ 8-bit Adam

  • จัดเก็บพารามิเตอร์ใน FP8; ปรับปรุงค่าใน FP32
  • EMA: ค่าเฉลี่ยเคลื่อนที่แบบเอกซ์โพเนนเชียล
  • คำนวณ EMA ของ gradient และ gradient กำลังสอง
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

8-bit Adam ประหยัดหน่วยความจำอย่างไร?

แผนภาพแสดง gradient ของพารามิเตอร์, EMA ของ gradient และ EMA ของ gradient กำลังสอง

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

8-bit Adam ประหยัดหน่วยความจำอย่างไร?

แผนภาพแสดง gradient ของพารามิเตอร์, EMA ของ gradient และ EMA ของ gradient กำลังสอง

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

8-bit Adam ประหยัดหน่วยความจำอย่างไร?

แผนภาพแสดง gradient ของพารามิเตอร์, EMA ของ gradient และ EMA ของ gradient กำลังสอง

  • แต่ละช่องคือพารามิเตอร์ แต่ละสีคือสถานะ
  • หน่วยความจำต่อพารามิเตอร์ = 2 ไบต์ = 1 ไบต์ต่อสถานะ * 2 สถานะ
  • หน่วยความจำรวม = หน่วยความจำต่อพารามิเตอร์ (2 ไบต์) * จำนวนพารามิเตอร์
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

ประมาณการใช้หน่วยความจำของ 8-bit Adam

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
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

ตั้งค่า 8-bit Adam Optimizer

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 ป้องกัน overfitting
  • decay_parameters: ระบุพารามิเตอร์ที่ใช้ weight decay
  • get_parameter_names: คืนชื่อพารามิเตอร์; ละเว้น layer nn.LayerNorm
  • ลบพารามิเตอร์ bias ออกจาก decay_parameters
  • ไม่ใช้ weight decay กับ normalization layer และ bias
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

ตั้งค่า 8-bit Adam Optimizer

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: กลุ่มหนึ่งใช้ weight decay; อีกกลุ่มไม่ใช้
  • beta1, beta2: อัตราการลดของ moment ที่ 1 และ 2; ค่ายิ่งสูง ยิ่งเสถียรแต่ช้า
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

ใช้งาน 8-bit Adam ด้วย 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}
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

ใช้งาน 8-bit Adam ด้วย 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
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

คำนวณการใช้หน่วยความจำของ 8-bit Adam

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
  • เปรียบเทียบกับ AdamW: 8-bit Adam ใช้หน่วยความจำเพียง 1/4
Number of AdamW parameters: 131,566,188
AdamW size: 502 MB
การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

มาฝึกกันเลย!

การเทรน AI Model อย่างมีประสิทธิภาพด้วย PyTorch

Preparing Video For Download...