Pelatihan hemat memori dengan Adafactor

Pelatihan Model AI Efisien dengan PyTorch

Dennis Lee

Data Engineer, Amazon

Optimizer untuk efisiensi pelatihan

 

 

Ikon yang merepresentasikan AdamW, Adafactor, dan Adam 8-bit.

Pelatihan Model AI Efisien dengan PyTorch

Kompromi optimizer

Diagram yang menunjukkan kompromi antara jumlah parameter dan presisi untuk AdamW, Adafactor, dan Adam 8-bit.

Pelatihan Model AI Efisien dengan PyTorch

Bagaimana Adafactor bekerja?

 

Diagram yang menampilkan langkah-langkah Adafactor.

Pelatihan Model AI Efisien dengan PyTorch

Bagaimana Adafactor bekerja?

 

Diagram yang menampilkan langkah-langkah Adafactor.

Pelatihan Model AI Efisien dengan PyTorch

Bagaimana Adafactor bekerja?

 

Diagram yang menampilkan langkah-langkah Adafactor.

 

  • EMA: exponential moving average (rata-rata bergerak eksponensial)
  • Momen kedua: EMA dari gradien kuadrat
Pelatihan Model AI Efisien dengan PyTorch

Bagaimana Adafactor bekerja?

 

Diagram yang menampilkan langkah-langkah Adafactor.

 

  • EMA: exponential moving average (rata-rata bergerak eksponensial)
  • Momen kedua: EMA dari gradien kuadrat
Pelatihan Model AI Efisien dengan PyTorch

Bagaimana Adafactor menghemat memori?

 

Diagram yang menampilkan matriks momen kedua, jumlah kolom, dan jumlah baris.

  • Hemat memori dengan tidak menyimpan matriks momen kedua
Pelatihan Model AI Efisien dengan PyTorch

Bagaimana Adafactor menghemat memori?

 

Diagram yang menampilkan matriks momen kedua, jumlah kolom, dan jumlah baris.

  • Hemat memori dengan tidak menyimpan matriks momen kedua
  • Sebagai gantinya, simpan jumlah kolom dan jumlah baris matriks
Pelatihan Model AI Efisien dengan PyTorch

Bagaimana Adafactor menghemat memori?

 

Diagram yang menampilkan matriks momen kedua, jumlah kolom, dan jumlah baris.

  • Hemat memori dengan tidak menyimpan matriks momen kedua
  • Sebagai gantinya, simpan jumlah kolom dan jumlah baris matriks
Pelatihan Model AI Efisien dengan PyTorch

Bagaimana Adafactor menghemat memori?

 

Diagram yang menampilkan matriks momen kedua, jumlah kolom, dan jumlah baris.

  • Hemat memori dengan tidak menyimpan matriks momen kedua
  • Sebagai gantinya, simpan jumlah kolom dan jumlah baris matriks
  • Perkirakan matriks penuh dengan mengalikan jumlah kolom dan jumlah baris
Pelatihan Model AI Efisien dengan PyTorch

Implementasi Trainer dan Accelerator

Bagan yang membandingkan kemudahan penggunaan vs. kemampuan kustomisasi untuk Accelerator dan Trainer.

Pelatihan Model AI Efisien dengan PyTorch

Implementasi Adafactor dengan Trainer

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

optim="adafactor")
trainer = Trainer(model=model, args=training_args, train_dataset=train_dataset, eval_dataset=validation_dataset, compute_metrics=compute_metrics) trainer.train()
{'epoch': 1.0, 'eval_accuracy': 0.6, 'eval_f1': 0.5}
Pelatihan Model AI Efisien dengan PyTorch

Implementasi Adafactor dengan Accelerator

# Mengasumsikan PyTorch 2.5 atau lebih baru
from torch.optim import Adafactor

optimizer = Adafactor(params=model.parameters(), lr=lr)
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.71
Pelatihan Model AI Efisien dengan PyTorch

Inspeksi state optimizer

  • Akses optimizer melalui state-nya
optimizer_state = optimizer.state.values()
  • Atau akses optimizer melalui trainer
optimizer_state = trainer.optimizer.state.values()
print(optimizer_state)
dict_values([{'step': tensor(3.),
              'exp_avg_sq_row': tensor([1.0000e-30, 1.0000e-30, 1.0000e-30,  ...]), 
              'exp_avg_sq_col': tensor([4.6147e-11, 5.5115e-11, 1.6338e-10, ...])}, ...])
Pelatihan Model AI Efisien dengan PyTorch

Hitung penggunaan memori Adafactor

total_size_megabytes, total_num_elements = \
    compute_optimizer_size(trainer.optimizer.state.values())
print(f"Jumlah parameter Adafactor: {total_num_elements:,}")
print(f"Ukuran Adafactor: {total_size_megabytes:.0f} MB")
Jumlah parameter Adafactor: 178,712
Ukuran Adafactor: 1 MB
  • Bandingkan dengan AdamW: Adafactor jauh lebih hemat memori!
Jumlah parameter AdamW: 131,566,188
Ukuran AdamW: 502 MB
Pelatihan Model AI Efisien dengan PyTorch

Ayo berlatih!

Pelatihan Model AI Efisien dengan PyTorch

Preparing Video For Download...