FairScale이란 무엇인가

FairScale이란 무엇인가

FairScale은 PyTorch 확장 라이브러리로, 고성능·대규모 학습을 지원해요. 기본 PyTorch 능력을 확장하면서 최신 SOTA 스케일링 기법을 추가해요. 최신 분산 학습 기법을 조합 가능한 모듈과 사용하기 쉬운 API로 제공해서, 제한된 자원으로 모델을 확장하려는 연구자의 도구함에서 기본적인 부분을 담당해요.

출처: FairScale What is FairScale (공식)

기술 범주

FairScale의 기능은 크게 세 범주로 나뉘어요.

  1. 병렬화(Parallelism) — 레이어 병렬화와 텐서 병렬화로 모델을 확장하는 기법이에요.
  2. 샤딩 메서드(Sharding Methods) — 메모리와 계산은 대개 트레이드오프 관계예요. 이 범주에서는 모델 레이어·파라미터, 옵티마이저 상태, 그래디언트를 샤딩해서 낮은 메모리 사용과 효율적인 계산을 동시에 달성하려고 해요.
  3. 최적화(Optimization) — 모델 규모와 무관하게 메모리 사용을 최적화하고, 하이퍼파라미터 튜닝 없이 학습하며, 학습 성능을 어떤 방식으로든 최적화하는 모든 기법을 다뤄요.

설계 원칙

FairScale API는 세 가지 원칙 아래 설계됐어요. 최소한의 인지 부하로 이해하고 사용할 수 있어야 하고(사용성), 여러 API를 학습 루프 안에서 매끄럽게 조합할 수 있어야 하며(모듈성), 스케일링·효율 면에서 최고 성능을 제공해야 해요(성능).

왜 필요한가

모델이 커지면 개별 GPU 메모리로는 감당할 수 없게 돼요. 여러 GPU에 걸쳐 모델을 나누어 담는 병렬화, 옵티마이저 상태나 그래디언트를 샤딩해 메모리를 아끼는 기법, 그리고 훈련 자체를 최적화하는 기법이 필요해져요. FairScale은 이런 기법들을 하나의 라이브러리에서 조합할 수 있게 해줘요.

더 알아보기