FairScale
FairScale
FairScale은 고성능·대규모 학습을 위한 PyTorch 확장 라이브러리예요. 기본 PyTorch 기능을 확장하면서 최신 SOTA 스케일링 기법을 더해 줘요. 최신 분산 학습 기법을 조합 가능한 모듈과 사용하기 쉬운 API 형태로 제공해서, 제한된 자원으로 모델을 확장하려는 연구자에게 유용해요.
설계 원칙
FairScale은 세 가지 원칙을 따르도록 설계됐어요.
- 사용성(Usability) — 인지 부하를 최소로 하면서 FairScale API를 이해하고 사용할 수 있어야 한다.
- 모듈성(Modularity) — 여러 FairScale API를 학습 루프 안에서 매끄럽게 조합할 수 있어야 한다.
- 성능(Performance) — FairScale API는 스케일링과 효율 측면에서 최고의 성능을 제공한다.
기술 범주
FairScale의 기능은 세 가지 범주로 나뉘어요.
- 병렬화(Parallelism) — 레이어 병렬화와 텐서 병렬화로 모델을 확장하는 기법이에요.
- 샤딩 메서드(Sharding Methods) — 모델 레이어·파라미터, 옵티마이저 상태, 그래디언트를 샤딩해서 낮은 메모리 사용과 효율적인 계산을 동시에 달성하려는 기법이에요.
- 최적화(Optimization) — 모델 규모와 무관하게 메모리 사용을 최적화하고, 하이퍼파라미터 튜닝 없이 학습하며, 성능을 최적화하는 모든 기법이에요.
FSDP (FullyShardedDataParallel)
FairScale의 FullyShardedDataParallel(FSDP)은 큰 NN 모델로 확장하는 데 권장되는 방법이에요. 이 라이브러리는 PyTorch로 업스트림되어, PyTorch가 공식적으로 Fully Sharded Data Parallel API를 제공하게 되었어요. 여기 있는 FSDP 버전은 과거 레퍼런스와 실험용으로 유지되고 있어요.
FSDP는 모델 파라미터, 그래디언트, 옵티마이저 상태를 데이터 병렬 프로세스들에 걸쳐 샤딩해서, 기존 DDP(DataParallel·DistributedDataParallel)보다 훨씬 큰 모델을 제한된 메모리로 학습할 수 있게 해줘요.