FairScale 샤딩과 옵티마이저
FairScale 샤딩과 옵티마이저 (OSS)
FairScale은 옵티마이저 상태를 샤딩해 메모리를 아끼는 기법도 제공해요. Optimizer State Sharding(OSS)은 옵티마이저 상태를 데이터 병렬 워커들에 걸쳐 분산해서, 큰 모델에서 옵티마이저가 차지하던 메모리 부담을 줄여요. 이는 DeepSpeed의 ZeRO 옵티마이저와 같은 아이디어 계열이에요.
샤딩의 목적
FairScale의 샤딩 메서드는 낮은 메모리 사용과 효율적인 계산을 동시에 달성하려는 범주예요. 모델 레이어·파라미터뿐 아니라 옵티마이저 상태와 그래디언트까지 샤딩 대상에 포함돼요.
Adam 같은 옵티마이저는 파라미터마다 모멘텀·분산 같은 상태를 유지해요. 이 상태는 종종 모델 파라미터 자체보다 큰 메모리를 차지해요. OSS는 이 옵티마이저 상태를 여러 워커에 나눠 담아, 전체 학습 메모리를 크게 줄여요.
샤딩 기법의 활용
샤딩 기법을 직접 쓰는 것 외에도, FairScale API들은 조합해서 쓸 수 있도록 설계됐어요. 예를 들어 FSDP로 파라미터를 샤딩하고, 옵티마이저 상태는 OSS 방식으로 분산하는 식으로 여러 기법을 학습 루프에서 함께 사용할 수 있어요.
내부 구조 (model_parallel)
FairScale의 nn.model_parallel 모듈은 Megatron-LM에서 포크된 것으로, NVIDIA CORPORATION이 저작권을 가진 Apache License 2.0 하에 배포돼요. 이 모듈을 통해 텐서 병렬화·레이어 병렬화 같은 기법을 쓸 수 있어요.