Hugging Face BLOOM 추론 최적화
Hugging Face는 일련의 반복적인 최적화를 통해 BLOOM 모델의 추론 지연 시간을 5배 줄이고 처리량을 50배 늘렸습니다. 최종 아키텍처는 Tensor Parallelism(TP)과 결합된 PyTorch, 어텐션을 위한 맞춤형 CUDA 커널, 그리고 커널 융합을 위한 torch.jit.script를 활용합니다.
파이프라인에서 텐서 병렬 처리로 전환
BLOOM(176B 파라미터, bf16 기준 352GB)의 초기 추론은 accelerate 라이브러리의 device_map="auto"를 사용한 파이프라인 병렬 처리(PP)로 구현되었습니다. PP에서는 각 GPU가 특정 레이어 집합을 소유하고 데이터를 순차적으로 처리한 뒤 다음 GPU에 전달합니다.
지연 시간을 줄이기 위해 Hugging Face는 텐서 병렬 처리(TP)로 전환했으며, TP에서는 각 GPU가 모든 레이어의 가중치 일부를 소유하여 모든 GPU가 동시에 작업할 수 있습니다. 이 전환으로 성능이 크게 향상되었습니다:
- Latency(지연 시간): 300ms/토큰에서 91ms/토큰으로 감소했습니다.
- Throughput(처리량): 초당 10 요청(RPS)으로 증가했습니다.
TP는 ncclAllReduce를 통한 통신 오버헤드를 도입하지만, 배치 처리(배치 크기 1과 32가 종종 비슷한 지연 시간을 보임)가 가능해져 전체 처리량이 크게 향상되었습니다.
PyTorch 기반 최적화
병렬 처리 전략 외에도, TensorBoard를 사용한 프로파일링을 통해 식별된 병목 현상을 제거하기 위해 여러 저수준 PyTorch 최적화가 구현되었습니다.
torch.jit.script를 이용한 커널 융합
Gelu 연산자는 원래 여러 개의 요소별 커널을 실행하여 과도한 텐서 복사를 일으켰습니다. bloom_gelu_forward 함수에 @torch.jit.script를 적용함으로써 Hugging Face는 이를 단일 커널 연산으로 융합했으며, 지연 시간이 91ms/토큰에서 81ms/토큰으로 감소했습니다.
효율적인 PyTorch 구현
- ALiBi 최적화: 위치 임베딩이 이전에 너무 많은 위치에서 계산되었습니다. 이 계산을 중앙 집중화함으로써 해당 연산이 10배 빨라졌습니다.
- 텐서 복사 감소: 프로파일링 결과 어텐션 경로가
reshape와transpose연산으로 인해 크게 부담을 받고 있었습니다. 가중치와 KV 캐시(“past”)를 재구성함으로써 이러한 불필요한 복사를 제거했습니다.
맞춤형 CUDA 커널 및 하드웨어 가속
torch.jit.script만으로는 충분하지 않은 핵심 경로를 추가로 최적화하기 위해, Hugging Face는 마스크된 fill과 softmax 연산을 융합하는 맞춤형 CUDA 커널을 개발했습니다.
구체적으로, 커널은 다음 순서를 최적화합니다:
- 어텐션 마스크를 사용하여 어텐션 점수에
masked_fill_적용. - 안정성을 위해 float32로
softmax계산.
커널 내부에서 필요한 합계와 누적에만 업캐스팅을 제한함으로써 지연 시간이 81ms/토큰에서 71ms/토큰으로 추가 감소했습니다.
웹 서버 아키텍처 및 요청 처리
다양한 파라미터와 길이를 가진 사용자 요청을 처리하기 위해, Hugging Face는 유연한 배치 시스템을 구현했습니다:
- Inter-process Communication(프로세스 간 통신):
torch.distributed가 별도 프로세스를 필요로 하기 때문에, 서버는 Redis pub/sub를 사용해 원시 문자열을 모든 프로세스로 배포합니다. - Custom Generation Loop(맞춤형 생성 루프): 표준
generate함수를 배치 내 각 멤버에 서로 다른 파라미터(예: 샘플링, top-p)를 적용하는 맞춤 구현으로 교체했습니다. - Dynamic Batch Extraction(동적 배치 추출): 동일 배치 내 긴 요청 때문에 짧은 요청이 지연되는 것을 방지하기 위해, 서버는 전체 배치가 끝날 때까지 기다리는 대신 토큰 제한에 도달하는 즉시 완료된 요청을 추출하고 반환합니다.
평가했지만 제외된 접근법
최적화 과정 전반에 걸쳐 여러 다른 경로가 탐색되었습니다:
- JAX/Flax on TPUs: 병렬 구현은 쉬웠지만, 팀은 Ray와 TPU 워커 통신에서 심각한 안정성 문제를 겪었으며, 컴파일에 대한 세밀한 제어가 부족했습니다.
- DeepSpeed: 최종 반복과 유사한 인상적인 결과를 제공했지만, 스트레스 상황에서 정기적인 커널 충돌(CUDA illegal access) 등 안정성 문제를 겪었습니다.
- Rust Implementation:
tch-rs를 사용해 더 나은 동시성 제어를 위해 Rust로 구현된 버전이 있었지만, 실제 성능 향상은 PyTorch 벤치마크에서 프로파일러가 활성화된 상태였기 때문임이 밝혀졌습니다. - ONNX/TensorRT: 텍스트 생성 루프에 필요한 유연성과 로짓 계산을 위해 텐서를 GPU에 유지해야 하는 요구사항에 비해 너무 경직된 것으로 판단되었습니다.