모델 프루닝 기법 구현

PyTorch Lightning으로 만드는 확장 가능한 AI 모델

Sergiy Tkachuk

Director, GenAI Productivity

프루닝 사용 시점

$$

  • 📱 엣지/임베디드 시스템 배포에 유용

$$

  • ➕ 양자화와 결합해 효율을 더 높일 수 있음

$$

  • ⚡ 지연 시간 또는 모델 크기 축소가 중요할 때 사용
PyTorch Lightning으로 만드는 확장 가능한 AI 모델

모델 프루닝이란?

$$

  • 신경망에서 중요도가 낮은 연결을 제거합니다
  • 저장/연산이 효율적인 희소 모델을 만듭니다
  • 일반적 방법: L1 비구조화 프루닝

프루닝 예시

PyTorch Lightning으로 만드는 확장 가능한 AI 모델

모델 프루닝이란?

import torch.nn.utils.prune as prune

prune.l1_unstructured(model.fc, name="weight",
                      amount=0.4)

print(model.fc.weight.data)
tensor([[ 0.25, -0.13,  0.05,  0.70],
        [-0.88,  0.31, -0.02,  0.44]]) # 프루닝 전


tensor([[ 0.25, -0.13,  0.00,  0.70],
        [ 0.00,  0.31,  0.00,  0.44]]) # 프루닝 후 (가중치 40%를 0으로 설정)
PyTorch Lightning으로 만드는 확장 가능한 AI 모델

프루닝 마스크 이해

$$

  • 프루닝은 대상 가중치 텐서마다 이진 마스크를 추가합니다.

$$

  • Mask = 1 → 가중치 유지
  • Mask = 0 → 순전파 시 가중치를 0으로 설정

$$

  • 마스크를 제거하기 전까지 가중치는 메모리에 그대로 저장됩니다.
PyTorch Lightning으로 만드는 확장 가능한 AI 모델

프루닝을 영구 적용하기

  • 기본적으로 프루닝된 가중치는 원래 텐서에 남아 있습니다
  • 프루닝을 확정하려면 재매개변수를 제거합니다
  • 희소 레이어를 0으로 채운 표준 레이어로 변환합니다
Sequential(
  (fc): Linear(
    in_features=128, out_features=64,
    bias=True
    (weight): PrunedParam()
  )
) # prune.remove 전
import torch.nn.utils.prune as prune

prune.remove(model.fc, 'weight')

# Print model structure
print(model)

Sequential(
  (fc): Linear(in_features=128,
               out_features=64,
               bias=True)
) # prune.remove 후
PyTorch Lightning으로 만드는 확장 가능한 AI 모델

프루닝 영향 평가

  • 원본 모델과 프루닝 모델의 성능을 비교합니다
  • 정확도는 소폭 하락, 크기/메모리는 크게 절감 기대
  • 배포 시 이 트레이드오프 수용 가능성 평가에 도움

프루닝 트레이드오프

PyTorch Lightning으로 만드는 확장 가능한 AI 모델

Ayo berlatih!

PyTorch Lightning으로 만드는 확장 가능한 AI 모델

Preparing Video For Download...