본문으로 건너뛰기

[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.cucute_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 모델의 성능 향상에 크게 기여할 것으로 기대됩니다.

참고 자료

⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.

댓글

관련 포스트

PR Analysis 의 다른글