그래디언트 체크포인팅과 로컬 SGD

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

Dennis Lee

Data Engineer, Amazon

학습 효율성 개선

 

 

메모리 효율성, 통신 효율성, 연산 효율성을 나타내는 아이콘

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

그래디언트 체크포인팅으로 메모리 효율성 개선

 

 

메모리 효율성, 통신 효율성, 연산 효율성을 나타내는 아이콘

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

로컬 SGD로 통신 효율성 해결

 

 

메모리 효율성, 통신 효율성, 연산 효율성을 나타내는 아이콘

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

그래디언트 체크포인팅이란?

  • 그래디언트 체크포인팅: 저장할 활성화 값을 선택해 메모리 절약
  • 예시: A + B = C 계산

노드와 엣지로 그래디언트 체크포인팅을 나타낸 그래프

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

그래디언트 체크포인팅이란?

  • 그래디언트 체크포인팅: 저장할 활성화 값을 선택해 메모리 절약
  • 예시: A + B = C 계산
    • 먼저 A, B를 계산한 후 C를 계산

노드와 엣지로 그래디언트 체크포인팅을 나타낸 그래프

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

그래디언트 체크포인팅이란?

  • 그래디언트 체크포인팅: 저장할 활성화 값을 선택해 메모리 절약
  • 예시: A + B = C 계산
    • 먼저 A, B를 계산한 후 C를 계산
    • 이후 순전파에 A, B 불필요
  • A, B를 저장해야 할까요, 제거해야 할까요?

노드와 엣지로 그래디언트 체크포인팅을 나타낸 그래프

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

그래디언트 체크포인팅이란?

  • 그래디언트 체크포인팅: 저장할 활성화 값을 선택해 메모리 절약
  • 예시: A + B = C 계산
    • 먼저 A, B를 계산한 후 C를 계산
    • 이후 순전파에 A, B 불필요
  • A, B를 저장해야 할까요, 제거해야 할까요?
    • 체크포인팅 미사용: A, B 저장

노드와 엣지로 그래디언트 체크포인팅을 나타낸 그래프

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

그래디언트 체크포인팅이란?

  • 그래디언트 체크포인팅: 저장할 활성화 값을 선택해 메모리 절약
  • 예시: A + B = C 계산
    • 먼저 A, B를 계산한 후 C를 계산
    • 이후 순전파에 A, B 불필요
  • A, B를 저장해야 할까요, 제거해야 할까요?
    • 체크포인팅 미사용: A, B 저장
    • 그래디언트 체크포인팅: A, B 제거

노드와 엣지로 그래디언트 체크포인팅을 나타낸 그래프

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

그래디언트 체크포인팅이란?

  • 그래디언트 체크포인팅: 저장할 활성화 값을 선택해 메모리 절약
  • 예시: A + B = C 계산
    • 먼저 A, B를 계산한 후 C를 계산
    • 이후 순전파에 A, B 불필요
  • A, B를 저장해야 할까요, 제거해야 할까요?
    • 체크포인팅 미사용: A, B 저장
    • 그래디언트 체크포인팅: A, B 제거
    • 역전파 시 A, B 재계산

노드와 엣지로 그래디언트 체크포인팅을 나타낸 그래프

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

그래디언트 체크포인팅이란?

  • 그래디언트 체크포인팅: 저장할 활성화 값을 선택해 메모리 절약
  • 예시: A + B = C 계산
    • 먼저 A, B를 계산한 후 C를 계산
    • 이후 순전파에 A, B 불필요
  • A, B를 저장해야 할까요, 제거해야 할까요?
    • 체크포인팅 미사용: A, B 저장
    • 그래디언트 체크포인팅: A, B 제거
    • 역전파 시 A, B 재계산
    • B 재계산 비용이 크면 저장

노드와 엣지로 그래디언트 체크포인팅을 나타낸 그래프

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

Trainer와 Accelerator

Accelerator와 Trainer의 사용 편의성 및 커스터마이징 능력을 비교한 차트

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

Trainer와 Accelerator

Accelerator와 Trainer의 사용 편의성 및 커스터마이징 능력을 비교한 차트

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

Trainer로 그래디언트 체크포인팅 사용

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







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

Trainer로 그래디언트 체크포인팅 사용

training_args = TrainingArguments(output_dir="./results",
                                  evaluation_strategy="epoch",
                                  gradient_accumulation_steps=4,
                                  gradient_checkpointing=True)

trainer = Trainer(model=model, args=training_args, train_dataset=dataset["train"], eval_dataset=dataset["validation"], compute_metrics=compute_metrics)
trainer.train()
{'epoch': 1.0, 'eval_loss': 0.73, 'eval_accuracy': 0.03, 'eval_f1': 0.05}
PyTorch로 AI 모델 효율적으로 학습시키기

Trainer에서 Accelerator로

Accelerator와 Trainer의 사용 편의성 및 커스터마이징 능력을 비교한 차트

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

Accelerator로 그래디언트 체크포인팅 사용

accelerator = Accelerator(gradient_accumulation_steps=2)


for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        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()
PyTorch로 AI 모델 효율적으로 학습시키기

Accelerator로 그래디언트 체크포인팅 사용

accelerator = Accelerator(gradient_accumulation_steps=2)
model.gradient_checkpointing_enable()

for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        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()
PyTorch로 AI 모델 효율적으로 학습시키기

로컬 SGD로 통신 효율성 개선

 

 

메모리 효율성, 통신 효율성, 연산 효율성을 나타내는 아이콘

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

로컬 SGD란?

로컬 SGD가 일정 스텝 후 그래디언트를 동기화하는 방식을 보여주는 다이어그램

  • 각 디바이스가 병렬로 그래디언트 계산
PyTorch로 AI 모델 효율적으로 학습시키기

로컬 SGD란?

로컬 SGD가 일정 스텝 후 그래디언트를 동기화하는 방식을 보여주는 다이어그램

  • 각 디바이스가 병렬로 그래디언트 계산
  • 그래디언트 동기화: 드라이버 노드가 각 디바이스의 모델 파라미터 업데이트
  • 로컬 SGD: 그래디언트 동기화 빈도 감소
PyTorch로 AI 모델 효율적으로 학습시키기

Accelerator로 로컬 SGD 사용





for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        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()

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

Accelerator로 로컬 SGD 사용

from accelerate.local_sgd import LocalSGD

with LocalSGD(accelerator=accelerator, model=model, local_sgd_steps=8, 
              enabled=True) as local_sgd:
    for index, batch in enumerate(dataloader):
        with accelerator.accumulate(model):
            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()

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

Accelerator로 로컬 SGD 사용

from accelerate.local_sgd import LocalSGD

with LocalSGD(accelerator=accelerator, model=model, local_sgd_steps=8, 
              enabled=True) as local_sgd:
    for index, batch in enumerate(dataloader):
        with accelerator.accumulate(model):
            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()
            local_sgd.step()
PyTorch로 AI 모델 효율적으로 학습시키기

연습해 봅시다!

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

Preparing Video For Download...