[onnxruntime] ONNX Runtime MoE 최적화: QMoE CPU GEMM 성능 대폭 향상
PR 링크: microsoft/onnxruntime#32668 상태: Merged | 변경: +210 / -6
들어가며
최근 대규모 언어 모델(LLM) 분야에서는 Mixture-of-Experts (MoE) 아키텍처가 큰 주목을 받고 있습니다. MoE는 모델의 파라미터 수를 늘리면서도 추론 시에는 일부 전문가(expert)만 활성화하여 연산량을 효율적으로 관리할 수 있다는 장점이 있습니다. 하지만 MoE 모델을 CPU 환경에서 효율적으로 추론하는 것은 여전히 도전적인 과제입니다. 특히, 각 전문가 내에서 수행되는 대규모 행렬 곱셈(GEMM) 연산은 병목 현상의 주된 원인이 됩니다.
Microsoft의 ONNX Runtime은 이러한 문제를 해결하기 위해 지속적으로 성능 최적화를 진행하고 있으며, 이번 PR(#32644 후속 작업)은 특히 양자화된 MoE(QMoE) 모델의 CPU 추론 성능을 획기적으로 개선하는 데 초점을 맞추고 있습니다. 이 PR은 기존의 전문가별 GEMM 호출 방식이 가진 오버헤드를 줄이고, MLAS(Microsoft Linear Algebra Subprograms) 라이브러리의 기능을 활용하여 여러 전문가의 연산을 효율적으로 배치(batch) 처리함으로써 성능을 크게 향상시켰습니다.
본 글에서는 이 PR에서 어떤 변경이 이루어졌고, 왜 이러한 변경이 성능 향상으로 이어졌는지, 그리고 이 최적화가 주는 일반적인 교훈은 무엇인지 코드 변경사항과 함께 자세히 분석해보겠습니다.
코드 분석
이번 PR의 핵심 변경사항은 onnxruntime/contrib_ops/cpu/moe/moe_quantization_cpu.cc 파일에 집중되어 있습니다. 주요 목표는 QMoE 모델에서 활성화된 전문가들의 GEMM 연산을 효율적으로 묶어(batching) 처리함으로써, 각 전문가별로 발생하는 스레드 풀 장벽(thread-pool barrier) 오버헤드를 줄이는 것입니다.
1. 기존 방식의 문제점: 과도한 스레드 풀 장벽
기존 방식에서는 활성화된 각 전문가에 대해 별도의 MlasQNBitGemmBatch 호출이 발생했습니다. 이는 특히 추론 시(decode)처럼 활성화되는 전문가 수가 많을 때(top_k 전문가), 각 GEMM 연산마다 스레드 풀의 동기화 및 장벽 동기화 오버헤드가 누적되어 성능 저하를 유발했습니다. PR 설명에 따르면, 이로 인해 레이어당 2 * num_active_experts 만큼의 스레드 풀 장벽이 발생했습니다.
// 기존 방식 (개념적 설명, 실제 diff와는 약간 다를 수 있음)
// ...
if (qnbit_fc1_.packed != nullptr && qnbit_fc2_.packed != nullptr && num_active_experts < max_expert_threads) {
// 각 활성 전문가에 대해 개별적으로 MlasQNBitGemmBatch 호출
for (int64_t i = 0; i < num_experts; ++i) {
if (expert_is_active[i]) {
// FC1 GEMM 호출
MlasQNBitGemmBatch(...);
// FC2 GEMM 호출
MlasQNBitGemmBatch(...);
}
}
}
// ...
2. 새로운 방식: 전문가 그룹핑 및 배치 처리
이번 PR에서는 이러한 비효율성을 개선하기 위해, 활성화된 전문가들을 토큰(token) 수를 기준으로 그룹화하고, 각 그룹별로 단 한 번의 MlasQNBitGemmBatch 호출을 수행하도록 변경했습니다. 이는 특히 추론 시나리오(각 토큰이 하나의 전문가를 활성화하는 경우)에서 효과적이며, 2 * top_k 번의 개별 호출을 단 2번의 배치 호출로 줄일 수 있습니다.
이 방식은 전문가 루프가 이미 직렬화된 경우(즉, 활성화된 전문가 수가 스레드 수보다 적은 경우)에만 적용됩니다. 활성화된 전문가 수가 스레드 수보다 많거나 같은 경우에는 기존의 전문가별 스레드 할당 방식이 더 빠르기 때문에 해당 경로를 그대로 유지합니다.
// 변경된 방식의 핵심 로직 (diff 발췌 및 요약)
// ...
int num_expert_threads = std::max(1, std::min(num_active_experts, max_expert_threads));
// ... (기존 전문가별 스레드 할당 로직)
// 새로운 그룹핑 및 배치 처리 로직
bool use_grouped_qnbit = use_qnbit_fc1 && use_qnbit_fc2 && num_expert_threads == 1 && activation_type_ == ActivationType::SwiGLU;
if (use_grouped_qnbit) {
// ... (전문가들을 토큰 수 기준으로 그룹화)
// 그룹별로 단일 MlasQNBitGemmBatch 호출
for (const auto& bucket : buckets) {
// ... (현재 버킷에 속한 전문가들에 대한 GEMM 파라미터 설정)
MlasQNBitGemmBatch<float>(
rows, n, k, count, qnbit_bits, qnbit_blk, qnbit_compute_type_,
gemm_params.data(), grouped_ws.get(), tp,
&mlas_backend_kernel_selector_config_
);
}
// ... (SwiGLU 활성화 함수 적용 및 결과 취합)
}
// ...
주요 변경점 상세:
- 전문가 그룹핑:
expert_token_map을 사용하여 활성화된 전문가들을 토큰 수 기준으로 그룹화합니다. 동일한 토큰 수를 가진 전문가들은 연속적으로 배치되어 하나의MlasQNBitGemmBatch호출로 처리됩니다. - 메모리 관리: 그룹핑된 전문가 연산을 위해
a1_all,c1_all,a2_all,c2_all과 같은 임시 FP32 행렬들이 할당됩니다. 이는 각 그룹의 모든 라우팅된 행을 한 번에 처리하기 위함입니다. 리뷰어의 피드백에 따라, 과도한 메모리 사용을 방지하기 위해 총 라우팅된 행의 크기에 대한 제한(8 MiB)이 추가되었고, 이를 초과할 경우 기존의 전문가별 처리 방식으로 폴백(fallback)하도록 수정되었습니다. (tianleiwu리뷰 댓글 참조) - Bias 처리: FP32 bias는 그대로 사용하고, FP16 bias만 변환하여 사용합니다. 이는 불필요한 데이터 복사를 줄입니다.
- SwiGLU 활성화: FC1과 FC2 GEMM 사이에 SwiGLU 활성화 함수가 적용됩니다. 그룹핑된 방식에서는 모든 라우팅된 행에 대해 이 활성화 함수를 적용합니다.
3. 테스트 및 검증
- 기존
MoETest.*24개 테스트는 모두 통과했습니다. - 새로운
*_PerExpertDispatch및*_SingleThread테스트 케이스가 추가되어, intra-op 스레드 풀 크기를 조절하여 기존 전문가별 처리 방식과 새로운 그룹핑 처리 방식이 올바르게 동작하는지, 그리고 서로 비교 검증할 수 있도록 했습니다. - 실제 모델(
LFM2.5-8B-A1B int4)을 사용한 Greedy generation 테스트에서 이전과 동일한 토큰 ID와 텍스트 출력을 생성함을 확인했습니다. 이는 기능적 정확성을 보장합니다.
재현성 관련 참고사항:
그룹핑 시 전문가를 토큰 수 기준으로 정렬하기 때문에, 전문가별 누적 순서가 변경될 수 있습니다. 추론 결과는 일반적으로 비트 단위로 동일하지만, 부동 소수점 연산의 반올림 오차로 인해 매우 드물게 결과가 달라질 수 있습니다 (16개 토큰 처리 시 8번째 유효 숫자 수준에서 관찰됨).
왜 이게 좋은가?
이 PR의 핵심적인 개선은 스레드 풀 장벽(thread-pool barrier) 오버헤드 감소와 MLAS 커널의 배치 처리 능력 활용에 있습니다.
성능 향상
PR 설명에 제시된 벤치마크 결과는 이러한 개선의 효과를 명확하게 보여줍니다:
- 모델: LFM2.5-8B-A1B int4 (96코어 x64 호스트)
- 시나리오: 630 토큰/모드
| 메트릭 | 변경 전 (Before) | 변경 후 (After) |
|---|---|---|
| 토큰 생성 속도 | 36.5 ms/token (27.4 tok/s) | 30.3 ms/token (33.0 tok/s) |
| 프롬프트 처리 (16토큰) | 249 ms | 142 ms |
이는 토큰 생성 속도에서 약 1.2배, 프롬프트 처리 성능에서 약 1.75배의 상당한 성능 향상을 의미합니다.
단일 QMoE 레이어에서도 32개 스레드 사용 시 0.749ms에서 0.438ms로 개선되었습니다.
성능 향상 비율은 threads / active_experts 비율과 밀접한 관련이 있습니다. 즉, 사용 가능한 스레드 수에 비해 활성화되는 전문가 수가 적을수록, 제거되는 장벽 오버헤드의 영향이 커져 성능 향상 폭이 커집니다. 예를 들어, top-k=4 추론 시 32개 스레드 환경에서 1.78배의 성능 향상을 보였습니다.
일반적인 교훈
- 병렬 처리의 함정: 멀티스레딩 환경에서 각 작업(여기서는 각 전문가의 GEMM)마다 스레드 풀 장벽을 사용하는 것은 의도치 않은 병목 현상을 유발할 수 있습니다. 특히 작업 수가 많을 때 이 오버헤드는 무시할 수 없게 됩니다.
- 라이브러리 기능 활용: MLAS와 같은 고성능 선형대수 라이브러리는 종종 배치 처리(batching)와 같은 고급 기능을 제공합니다. 이러한 기능을 적극적으로 활용하면 개별적인 연산 호출을 줄이고 라이브러리 내부의 최적화된 병렬 처리 능력을 최대한 활용할 수 있습니다.
- 시나리오별 최적화: MoE 모델의 추론은 크게 두 가지 시나리오로 나뉩니다. 하나는 프롬프트 처리(prompt processing) 단계로, 여러 토큰이 여러 전문가에게 라우팅될 수 있습니다. 다른 하나는 토큰 생성(token generation) 단계로, 일반적으로 각 토큰은 하나의 전문가에게만 라우팅됩니다. 이 PR은 특히 후자의 시나리오에서 발생하는 오버헤드를 줄이는 데 효과적입니다. 최적화는 특정 시나리오에 맞춰 설계될 때 더 큰 효과를 발휘할 수 있습니다.
- 메모리 사용량과 성능의 균형: 새로운 배치 처리 방식은 성능을 향상시키지만, 더 많은 임시 메모리를 사용할 수 있습니다. 리뷰어의 피드백처럼, 이러한 메모리 사용량을 제한하고 필요시 이전 방식으로 폴백하는 전략은 메모리 제약이 있는 환경에서도 최적화의 이점을 누릴 수 있게 해줍니다.
결론
이번 ONNX Runtime PR은 QMoE 모델의 CPU 추론 성능을 개선하기 위한 매우 효과적인 최적화를 성공적으로 구현했습니다. 활성화된 전문가들을 토큰 수 기준으로 그룹화하고 MLAS의 배치 처리 기능을 활용함으로써, 스레드 풀 장벽 오버헤드를 크게 줄이고 실제 벤치마크에서 상당한 성능 향상을 달성했습니다. 이는 MoE 모델을 CPU 환경에서 더 효율적으로 배포하고 활용할 수 있는 기반을 마련했다는 점에서 큰 의미가 있습니다.
참고 자료
- https://onnxruntime.ai/docs/reference/contrib_ops/moe_quantization.html
- https://github.com/microsoft/onnxruntime/pull/32644
- https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#moe-operators
- https://github.com/microsoft/onnxruntime/blob/main/docs/UsingONNXRuntime.md#cpu-execution-provider
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [onnxruntime] ONNX Runtime CUDA EP, 2-bit 양자화 GEMM/GEMV 지원 추가로 모델 경량화 가속
- [vllm] [ROCm 성능 최적화] vLLM의 Fused Shared-Expert Gate GEMM 경로 개선 분석
- [flashinfer] FlashInfer의 GEMM 성능 혁신: cuTile 백엔드 도입과 최적화 여정
- [flashinfer] FlashInfer SM120 MoE GEMM 최적화: 웨이브+잔여물 비용 모델 도입
- [flashinfer] FlashInfer: NVIDIA Blackwell(SM120)을 위한 고성능 FP8 MoE GEMM 최적화
PR Analysis 의 다른글
- 이전글 [sglang] SGLang의 새로운 캐시 전략: T-LRU로 에이전트 워크로드의 TTFT 최적화하기
- 현재글 : [onnxruntime] ONNX Runtime MoE 최적화: QMoE CPU GEMM 성능 대폭 향상
- 다음글 [onnxruntime] ONNX Runtime의 RISC-V RVV 커널 최적화: 추론 성능 극대화
댓글