Antrenament echilibrat cu AdamW

Antrenament eficient al modelelor AI cu PyTorch

Dennis Lee

Data Engineer, Amazon

Antrenament eficient

 

 

Diagramă cu subiectele capitolului, cu evidențierea optimizatorilor.

Antrenament eficient al modelelor AI cu PyTorch

Optimizatori pentru eficiența antrenamentului

 

 

Diagramă cu trei optimizatori: AdamW, Adafactor și 8-bit Adam.

Antrenament eficient al modelelor AI cu PyTorch

Optimizatori pentru eficiența antrenamentului

 

 

Diagramă cu trei optimizatori: AdamW, Adafactor și 8-bit Adam.

Antrenament eficient al modelelor AI cu PyTorch

Optimizatori pentru eficiența antrenamentului

 

 

Diagramă cu trei optimizatori: AdamW, Adafactor și 8-bit Adam.

Antrenament eficient al modelelor AI cu PyTorch

Compromisuri între optimizatori

Diagramă cu compromisurile dintre numărul de parametri și precizie pentru AdamW, Adafactor și 8-bit Adam.

Antrenament eficient al modelelor AI cu PyTorch

Compromisuri între optimizatori

Diagramă cu compromisurile dintre numărul de parametri și precizie pentru AdamW, Adafactor și 8-bit Adam.

Antrenament eficient al modelelor AI cu PyTorch

Compromisuri între optimizatori

Diagramă cu compromisurile dintre numărul de parametri și precizie pentru AdamW, Adafactor și 8-bit Adam.

Antrenament eficient al modelelor AI cu PyTorch

Compromisuri între optimizatori

Diagramă cu compromisurile dintre numărul de parametri și precizie pentru AdamW, Adafactor și 8-bit Adam.

Antrenament eficient al modelelor AI cu PyTorch

Cum funcționează AdamW?

Diagramă care ilustrează funcționarea AdamW.

Antrenament eficient al modelelor AI cu PyTorch

Cum funcționează AdamW?

Diagramă care ilustrează funcționarea AdamW.

Antrenament eficient al modelelor AI cu PyTorch

Cum funcționează AdamW?

Diagramă care ilustrează funcționarea AdamW.

  • Calculează media mobilă exponențială (EMA) a gradienților
Antrenament eficient al modelelor AI cu PyTorch

Cum funcționează AdamW?

Diagramă care ilustrează funcționarea AdamW.

  • Calculează media mobilă exponențială (EMA) a gradienților
  • Calculează EMA a gradienților la pătrat
Antrenament eficient al modelelor AI cu PyTorch

Cum funcționează AdamW?

Diagramă care ilustrează funcționarea AdamW.

  • Calculează media mobilă exponențială (EMA) a gradienților
  • Calculează EMA a gradienților la pătrat
Antrenament eficient al modelelor AI cu PyTorch

Cum funcționează AdamW?

Diagramă care ilustrează funcționarea AdamW.

  • Calculează media mobilă exponențială (EMA) a gradienților
  • Calculează EMA a gradienților la pătrat
Antrenament eficient al modelelor AI cu PyTorch

Utilizarea memoriei de AdamW

Diagramă cu cantitățile implicate în calculele AdamW: gradienții parametrilor, EMA a gradienților și EMA a gradienților la pătrat.

  • Fiecare pătrat reprezintă un parametru, iar fiecare culoare o stare
Antrenament eficient al modelelor AI cu PyTorch

Utilizarea memoriei de AdamW

Diagramă cu cantitățile implicate în calculele AdamW: gradienții parametrilor, EMA a gradienților și EMA a gradienților la pătrat.

  • Fiecare pătrat reprezintă un parametru, iar fiecare culoare o stare
Antrenament eficient al modelelor AI cu PyTorch

Utilizarea memoriei de AdamW

Diagramă cu cantitățile implicate în calculele AdamW: gradienții parametrilor, EMA a gradienților și EMA a gradienților la pătrat.

  • Fiecare pătrat reprezintă un parametru, iar fiecare culoare o stare
  • Memorie per parametru = 8 octeți = 4 octeți per stare * 2 stări
  • Memorie totală = Memorie per parametru (8 octeți) * Număr de parametri
Antrenament eficient al modelelor AI cu PyTorch

Estimarea utilizării memoriei de 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
Antrenament eficient al modelelor AI cu PyTorch

Trainer și Accelerator

Diagramă cu compromisul dintre posibilitatea de personalizare și ușurința de utilizare pentru Accelerator și Trainer.

Antrenament eficient al modelelor AI cu PyTorch

Implementarea AdamW cu 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}
Antrenament eficient al modelelor AI cu PyTorch

Implementarea AdamW cu 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
Antrenament eficient al modelelor AI cu PyTorch

Inspectarea stării optimizatorului

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.], ...])}, ...])
Antrenament eficient al modelelor AI cu PyTorch

Calcularea dimensiunii optimizatorului

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
Antrenament eficient al modelelor AI cu PyTorch

Calcularea dimensiunii optimizatorului

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
Antrenament eficient al modelelor AI cu PyTorch

Să exersăm!

Antrenament eficient al modelelor AI cu PyTorch

Preparing Video For Download...