[flashinfer] FlashInfer, MoE 모델의 성능을 극적으로 향상시키는 융합 커널과 최적화된 스케줄러 도입
PR 링크: flashinfer-ai/flashinfer#4130 상태: Merged | 변경: +3646 / -151
들어가며
최근 대규모 언어 모델(LLM) 분야에서는 Mixture-of-Experts (MoE) 아키텍처가 큰 주목을 받고 있습니다. MoE는 모델의 파라미터 수를 늘리면서도 추론 시 활성화되는 파라미터 수를 제한하여 효율성을 높이는 방식입니다. 하지만 MoE 모델의 성능을 최대한 끌어내기 위해서는 각 전문가(expert) 내에서의 연산, 특히 Feed-Forward Network (FFN)의 첫 번째 계층(FC1)에서의 행렬 곱셈(GEMM)과 활성화 함수 적용을 효율적으로 처리하는 것이 중요합니다.
이번 글에서는 FlashInfer 라이브러리의 최신 Pull Request(PR)에서 이루어진 두 가지 주요 개선 사항을 심층적으로 분석합니다. 첫째, MoE 모델의 FC1 연산을 단일 커널로 융합(fused)하여 성능을 향상시키는 새로운 기능이 추가되었습니다. 둘째, MoE 스케줄링 방식을 워프 협력(warp-cooperative) 방식으로 개선하여 GPU 활용률을 높였습니다. 이 PR은 특히 SM120 아키텍처를 타겟으로 하며, FP8 및 MXFP8 정밀도를 지원합니다.
코드 분석
이번 PR의 핵심 변경 사항은 csrc/cute_sm120_mxfp8_groupwise/ 디렉토리 내의 CUDA C++ 코드에서 주로 이루어졌습니다. 주요 변경점을 파일별로 살펴보겠습니다.
1. cute_sm120_fp8_op.cu 및 cute_sm120_fp8_op_jit_binding.cu:
이 파일들은 FP8 정밀도를 사용하는 MoE GEMM 연산을 정의하고 JIT 컴파일을 위한 바인딩을 제공합니다. 가장 중요한 변경은 CutlassFP8GroupwiseMoeGEMMSM120 함수의 시그니처에 is_gated 파라미터가 추가된 것입니다.
Before:
void CutlassFP8GroupwiseMoeGEMMSM120(TensorView a, TensorView b, TensorView a_scale,
TensorView b_scale, TensorView m_indptr, TensorView out,
std::string scale_major_mode, int64_t scale_granularity_m,
- int64_t scale_granularity_n, int64_t scale_granularity_k) {
+ int64_t scale_granularity_n, int64_t scale_granularity_k,
+ int64_t is_gated) {
After:
void CutlassFP8GroupwiseMoeGEMMSM120(TensorView a, TensorView b, TensorView a_scale,
TensorView b_scale, TensorView m_indptr, TensorView out,
std::string scale_major_mode, int64_t scale_granularity_m,
- int64_t scale_granularity_n, int64_t scale_granularity_k) {
+ int64_t scale_granularity_n, int64_t scale_granularity_k,
+ int64_t is_gated) {
is_gated 플래그는 새로운 융합 연산을 활성화하는 데 사용됩니다. 이 플래그가 활성화되면, 입력 가중치 b는 두 부분으로 나뉘어 저장됩니다. 앞부분([0, N))에는 업 프로젝션(up-projection) 가중치가, 뒷부분([N, 2*N))에는 게이트(gate) 프로젝션 가중치가 저장됩니다. 커널은 이 두 가중치를 사용하여 GEMM 연산을 수행한 후, SiLU 활성화 함수를 게이트에 적용하고 그 결과를 업 프로젝션 결과와 곱하는 SiLU(gate) * up 연산을 단일 커널 내에서 처리합니다. 이는 기존 방식에서 별도의 커널로 수행되던 GEMM과 활성화 함수 적용을 하나로 합쳐, 중간 결과의 전역 메모리 쓰기 및 읽기 작업을 제거함으로써 성능을 향상시킵니다.
또한, is_gated=True일 경우 출력의 너비(out.size(1))가 입력 가중치 b의 너비(n)의 절반(out_n = n / 2)이 되도록 변경되었습니다. 이는 융합된 연산의 결과로 최종 출력 차원이 줄어들기 때문입니다. 이에 따라 출력 차원에 대한 검증 로직도 다음과 같이 수정되었습니다.
Before:
- TVM_FFI_ICHECK_EQ(out.size(1), n)
- << "out.size(1) (" << out.size(1) << ") must match b.size(1) (" << n << ")";
+ TVM_FFI_ICHECK_EQ(out.size(1), out_n)
+ << "out.size(1) (" << out.size(1) << ") must match output N (" << out_n
+ << (gated ? "; = b.size(1)/2 for gated" : "; = b.size(1)") << ")";
After:
- TVM_FFI_ICHECK_EQ(out.size(1), n)
- << "out.size(1) (" << out.size(1) << ") must match b.size(1) (" << n << ")";
+ TVM_FFI_ICHECK_EQ(out.size(1), out_n)
+ << "out.size(1) (" << out.size(1) << ") must match output N (" << out_n
+ << (gated ? "; = b.size(1)/2 for gated" : "; = b.size(1)") << ")";
그리고 is_gated 모드에서는 출력 너비(out_n)가 16의 배수여야 한다는 제약 조건이 추가되었습니다.
Before:
- TVM_FFI_ICHECK_EQ(n % 16, 0) << "n must be multiple of 16; got n=" << n;
+ TVM_FFI_ICHECK_EQ(out_n % 16, 0) << "output N must be multiple of 16; got " << out_n;
After:
- TVM_FFI_ICHECK_EQ(n % 16, 0) << "n must be multiple of 16; got n=" << n;
+ TVM_FFI_ICHECK_EQ(out_n % 16, 0) << "output N must be multiple of 16; got " << out_n;
2. cute_sm120_fp8_runner.cu:
이 파일은 실제 커널 실행을 담당하는 CuteSm120Fp8GemmRunner 클래스의 구현을 포함합니다. moe_gemm_fp8_nt_groupwise 함수가 is_gated 파라미터를 받아, 해당 플래그에 따라 기존의 moe_gemm_fp8_nt_groupwise_impl을 호출하거나 새로운 fused_moe_fp8_nt_groupwise_impl을 호출하도록 수정되었습니다.
Before:
int scale_granularity_m, int scale_granularity_n,
int scale_granularity_k) {
check_scale_granularity_mnk(scale_granularity_m, scale_granularity_n, scale_granularity_k);
- moe_gemm_fp8_nt_groupwise_impl(D, A, B, token_offset, num_experts, total_rows, shape_n, shape_k,
- stream, SFA, SFB);
}
After:
int scale_granularity_m, int scale_granularity_n,
- int scale_granularity_k) {
+ int scale_granularity_k, bool is_gated) {
check_scale_granularity_mnk(scale_granularity_m, scale_granularity_n, scale_granularity_k);
- moe_gemm_fp8_nt_groupwise_impl(D, A, B, token_offset, num_experts, total_rows, shape_n, shape_k,
- stream, SFA, SFB);
+ if (is_gated) {
+ fused_moe_fp8_nt_groupwise_impl(D, A, B, token_offset, num_experts, total_rows, shape_n,
+ shape_k, stream, SFA, SFB);
+ } else {
+ moe_gemm_fp8_nt_groupwise_impl(D, A, B, token_offset, num_experts, total_rows, shape_n, shape_k,
+ stream, SFA, SFB);
+ }
}
fused_moe_fp8_nt_groupwise_impl 함수는 새로운 융합 커널 로직을 포함하며, select_fp8_fused_moe_tile_m 함수를 사용하여 최적의 타일 크기(tile_m)를 동적으로 선택합니다. 이 함수는 입력 total_rows, shape_n, num_experts, 그리고 사용 가능한 SM 수(num_sms)를 고려하여 가장 효율적인 tile_m 값을 결정합니다.
3. sm120_common/moe_scheduler.cuh (새로운 파일, diff에는 포함되지 않음):
이 파일은 새로운 워프 협력 MoE 스케줄러를 구현합니다. 이전의 선형 스캔 방식 대신, 워프 내에서 협력적으로 타일 인덱스를 계산하는 방식을 사용합니다. 이는 특히 ZeroPadding 토큰이 많은 경우에 스케줄링 오버헤드를 크게 줄여줍니다.
왜 이게 좋은가?
이번 PR은 두 가지 핵심적인 최적화를 통해 MoE 모델의 추론 성능을 크게 향상시켰습니다.
1. FC1 연산 융합 (Fused MoE FC1)
- 문제점: 기존에는 MoE 모델의 FC1 계층에서 가중치 행렬 곱셈(GEMM)과 활성화 함수(SiLU) 적용이 별도의 커널로 실행되었습니다. 이 과정에서 GEMM 결과가 전역 메모리에 쓰여지고, 다시 활성화 함수 커널에서 읽어오는 비효율이 발생했습니다. 특히 작은 배치 크기나 짧은 시퀀스 길이에서는 이 메모리 I/O 비용이 연산 자체의 비용보다 커질 수 있었습니다.
- 개선:
is_gated=True옵션을 통해 GEMM과 SiLU 활성화 함수를 단일 커널로 융합했습니다. 이로써SiLU(gate) * up연산이 GPU 레지스터 내에서 직접 처리될 수 있게 되어, 전역 메모리 접근이 불필요해졌습니다. 결과적으로 연산 속도가 향상되고 메모리 대역폭 사용량이 감소했습니다. - 성능 향상: 벤치마크 결과에 따르면, 이 융합은 특히 긴 시퀀스 길이(prefill 단계)에서 상당한 성능 향상을 보였습니다. 예를 들어, Qwen3.5-35B 모델에서 FP8 정밀도 사용 시 최대 33.1%의 속도 향상을 기록했습니다. RTX PRO 5000 Blackwell에서도 최대 23.9%의 향상을 보였습니다. 이러한 성능 향상은 모델의 크기(K 값)가 작을수록 더 두드러지는 경향을 보였습니다. 이는 작은 K 값에서는 메모리 I/O 병목 현상이 더 심각했음을 시사합니다.
2. 워프 협력 MoE 스케줄러
- 문제점: 이전의 MoE 스케줄러는 토큰-타일 매핑을 위해 각 워프(warp)가 독립적으로 선형 스캔을 수행했습니다. 이는 특히 ZeroPadding 토큰이 많은 경우, 불필요한 메모리 접근과 계산을 유발하여 스케줄링 오버헤드가 커지는 원인이 되었습니다.
- 개선: 새로운 스케줄러는 워프 내에서 협력적인 스캔 방식을 도입했습니다. 32개의 스레드로 구성된 워프가 데이터를 모아 로드하고,
shfl명령어를 이용한 접두사 합(prefix-sum) 계산 및ballot연산을 통해 효율적으로 타일 인덱스를 결정합니다. 이는 CUTLASS의 SM90 그룹 스케줄러와 유사한 방식입니다. - 성능 향상: 이 개선은 특히 디코드(decode) 시나리오에서 스케줄링 오버헤드를 크게 줄였습니다. RTX PRO 6000 Blackwell에서 마이크로벤치마크 결과, 선형 스캔 방식 대비 GEMM1 연산에서 76%, GEMM2 연산에서 146%의 속도 향상을 보였습니다. 이는 전역 메모리 로드 트래픽을 7.1배, 명령어 수를 5배 감소시킨 결과입니다. 스케줄링 오버헤드가 이상적인 값의 2~3배에서 약 20% 수준으로 감소했습니다.
일반적 교훈
- 융합의 힘: 연산과 활성화 함수를 융합하는 것은 GPU에서 메모리 I/O 병목을 줄이는 강력한 기법입니다. 특히 LLM과 같이 연산 집약적이면서도 메모리 접근 패턴이 예측 가능한 경우, 융합 커널은 상당한 성능 향상을 가져올 수 있습니다.
- 협력적 병렬 처리: 워프 내 스레드 간의 협력을 통해 스케줄링 및 동기화 오버헤드를 줄일 수 있습니다. 이는 GPU 아키텍처의 특성을 잘 활용하는 중요한 최적화 전략입니다.
- 정밀도와 성능: FP8 및 MXFP8과 같은 저정밀도 연산을 지원하는 것은 LLM 추론 성능 향상에 필수적입니다. 이번 PR은 이러한 저정밀도 연산에서의 융합 및 스케줄링 최적화를 성공적으로 보여줍니다.
리뷰 피드백 반영
리뷰 과정에서 is_gated 플래그의 동작 방식, 특히 입력 b 텐서가 업 프로젝션과 게이트 프로젝션을 어떻게 패킹하는지에 대한 명확성이 요구되었습니다. CarstyYou님의 지적에 따라, core.py 파일의 문서 문자열과 주석이 업데이트되어 b 텐서의 [0, N) 열에는 업 프로젝션 가중치가, [N, 2N) 열에는 게이트 프로젝션 가중치가 저장됨을 명확히 했습니다. 이는 커널(mB_up = mB, mB_gate = domain_offset(N))의 동작과 일치하도록 수정되었습니다.
또한, 테스트 실패 관련 피드백이 있었으나, 이는 주로 CI 환경 설정 문제로 보이며 코드 자체의 결함보다는 인프라 문제로 분류되었습니다. jiahanc님의 /bot run tests/moe 및 /bot run tests/grouped_mm 명령어는 해당 테스트들을 재실행하여 코드 변경의 정확성을 검증하는 데 사용되었습니다.
결론
이번 FlashInfer PR은 MoE 모델의 추론 성능을 한 단계 끌어올리는 중요한 개선을 이루었습니다. FC1 연산의 융합 커널 도입과 워프 협력 스케줄러 구현은 LLM의 효율적인 배포를 위한 핵심적인 기술 발전입니다. 특히 SM120 아키텍처와 FP8/MXFP8 정밀도에 최적화된 이 기술들은 향후 더 크고 복잡한 MoE 모델의 성능 향상에 크게 기여할 것으로 기대됩니다.
참고 자료
- https://github.com/flashinfer-ai/flashinfer/blob/main/csrc/cute_sm120_mxfp8_groupwise/cute_sm120_fp8_op.cu
- https://github.com/flashinfer-ai/flashinfer/blob/main/csrc/cute_sm120_mxfp8_groupwise/cute_sm120_fp8_runner.cu
- https://github.com/flashinfer-ai/flashinfer/blob/main/csrc/cute_sm120_mxfp8_groupwise/sm120_blockscaling/launch.cuh
- https://github.com/flashinfer-ai/flashinfer/blob/main/csrc/cute_sm120_mxfp8_groupwise/sm120_common/moe_scheduler.cuh
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
PR Analysis 의 다른글
- 이전글 [transformers] Hugging Face Transformers: NoRepeatNGramLogitsProcessor 벡터화 및 성능 최적화
- 현재글 : [flashinfer] FlashInfer, MoE 모델의 성능을 극적으로 향상시키는 융합 커널과 최적화된 스케줄러 도입
- 다음글 [vllm] vLLM 컴파일 최적화: Transformers 모델을 위한 FusedAddRMSNorm 도입
댓글