8비트 Adam을 활용한 혼합 정밀도 훈련

PyTorch로 AI 모델 효율적으로 학습시키기

Dennis Lee

Data Engineer, Amazon

훈련 효율성을 위한 옵티마이저

AdamW, Adafactor, 8비트 Adam의 파라미터 수와 정밀도 간의 트레이드오프를 보여주는 다이어그램.

PyTorch로 AI 모델 효율적으로 학습시키기

훈련 효율성을 위한 옵티마이저

AdamW, Adafactor, 8비트 Adam의 파라미터 수와 정밀도 간의 트레이드오프를 보여주는 다이어그램.

PyTorch로 AI 모델 효율적으로 학습시키기

8비트 Adam은 어떻게 작동하나요?

8비트 Adam의 단계를 보여주는 다이어그램.

  • 파라미터는 FP8로 저장하고 FP32로 최적화
PyTorch로 AI 모델 효율적으로 학습시키기

8비트 Adam은 어떻게 작동하나요?

8비트 Adam의 단계를 보여주는 다이어그램.

  • 파라미터는 FP8로 저장하고 FP32로 최적화
PyTorch로 AI 모델 효율적으로 학습시키기

8비트 Adam은 어떻게 작동하나요?

8비트 Adam의 단계를 보여주는 다이어그램.

  • 파라미터는 FP8로 저장하고 FP32로 최적화
PyTorch로 AI 모델 효율적으로 학습시키기

8비트 Adam은 어떻게 작동하나요?

8비트 Adam의 단계를 보여주는 다이어그램.

  • 파라미터는 FP8로 저장하고 FP32로 최적화
  • EMA: 지수 이동 평균
  • 그레이디언트 및 그레이디언트 제곱의 EMA 계산
PyTorch로 AI 모델 효율적으로 학습시키기

8비트 Adam은 어떻게 작동하나요?

8비트 Adam의 단계를 보여주는 다이어그램.

  • 파라미터는 FP8로 저장하고 FP32로 최적화
  • EMA: 지수 이동 평균
  • 그레이디언트 및 그레이디언트 제곱의 EMA 계산
PyTorch로 AI 모델 효율적으로 학습시키기

8비트 Adam은 어떻게 메모리를 절약하나요?

파라미터 그레이디언트, 그레이디언트의 EMA, 그레이디언트 제곱의 EMA를 나타내는 다이어그램.

PyTorch로 AI 모델 효율적으로 학습시키기

8비트 Adam은 어떻게 메모리를 절약하나요?

파라미터 그레이디언트, 그레이디언트의 EMA, 그레이디언트 제곱의 EMA를 나타내는 다이어그램.

PyTorch로 AI 모델 효율적으로 학습시키기

8비트 Adam은 어떻게 메모리를 절약하나요?

파라미터 그레이디언트, 그레이디언트의 EMA, 그레이디언트 제곱의 EMA를 나타내는 다이어그램.

  • 각 사각형은 파라미터, 각 색상은 상태를 나타냄
  • 파라미터당 메모리 = 2바이트 = 상태당 1바이트 × 2 상태
  • 총 메모리 = 파라미터당 메모리(2바이트) × 파라미터 수
PyTorch로 AI 모델 효율적으로 학습시키기

8비트 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
PyTorch로 AI 모델 효율적으로 학습시키기

8비트 Adam 옵티마이저 설정

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)는 과적합을 방지
  • decay_parameters: 가중치 감쇠를 적용할 파라미터 지정
  • get_parameter_names: 파라미터 이름 반환; nn.LayerNorm 레이어 제외
  • decay_parameters에서 bias 파라미터 제거
  • 정규화 레이어와 바이어스에는 가중치 감쇠 미적용
PyTorch로 AI 모델 효율적으로 학습시키기

8비트 Adam 옵티마이저 설정

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: 한 그룹은 가중치 감쇠 적용, 다른 그룹은 미적용
  • beta1, beta2: 1차 및 2차 모멘트의 감쇠율; 높을수록 안정적이나 훈련 속도 느림
PyTorch로 AI 모델 효율적으로 학습시키기

Trainer로 8비트 Adam 구현

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}
PyTorch로 AI 모델 효율적으로 학습시키기

Accelerator로 8비트 Adam 구현

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
PyTorch로 AI 모델 효율적으로 학습시키기

8비트 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비트 Adam은 메모리를 1/4만 사용
Number of AdamW parameters: 131,566,188
AdamW size: 502 MB
PyTorch로 AI 모델 효율적으로 학습시키기

연습해 봅시다!

PyTorch로 AI 모델 효율적으로 학습시키기

Preparing Video For Download...