Fused MoE 모듈러 커널

Fused MoE 모듈러 커널 (Modular Kernel)

MoE 모델을 서빙할 때 가장 핵심이 되는 연산 중 하나가 fused MoE 커널이에요. 그런데 정작 이를 구현하려고 보면, 입력 액티베이션 형태에 따라 다른 경로를 타야 하고 각 단계마다 다양한 구현이 가능해서 조합 수가 어마어마하게 불어나요. vLLM은 이 문제를 모듈러 커널(modular kernel) 프레임워크로 풀어요. 이 문서는 그 구조를 설명해요.

출처: 공식문서

소개

FusedMoEModularKernel의 실제 구현은 여기 있어요.

입력 액티베이션의 형태에 따라 fused MoE 구현은 크게 두 종류로 나뉘어요.

  • Contiguous / Standard / Non-Batched, 그리고
  • Batched

!!! note 이 문서에서 Contiguous, Standard, Non-Batched는 같은 의미로 서로 바꿔 쓸 수 있어요.

입력 액티베이션 형태는 전적으로 어떤 All2All Dispatch를 쓰느냐에 달려 있어요.

  • Contiguous 변형에서 All2All Dispatch는 액티베이션을 (M, K) 형태의 contiguous 텐서로, TopK Ids와 TopK weights를 (M, num_topk) 형태로 돌려줘요. DeepEPHTPrepareAndFinalize가 그 예시예요.
  • Batched 변형에서 All2All Dispatch는 액티베이션을 (num_experts, max_tokens, K) 형태로 돌려줘요. 여기서 같은 expert를 구독하는 액티베이션/토큰들이 한데 묶여요. 이때 텐서의 모든 항목이 유효한 건 아니에요. 보통 액티베이션 텐서와 함께 크기가 num_expertsexpert_num_tokens 텐서가 따라오는데, expert_num_tokens[i]는 i번째 expert를 구독하는 유효 토큰 수를 나타내요. DeepEPLLPrepareAndFinalize가 그 예시예요.

fused MoE 연산은 Contiguous와 Batched 변형 모두에서 일반적으로 여러 연산으로 구성돼요. 아래 다이어그램이 그 흐름을 보여줘요.

Fused MoE Non-Batched

Fused MoE Batched

!!! note 연산 구성 측면에서 Batched와 Non-Batched의 가장 큰 차이는 Permute / Unpermute 연산 유무예요. 나머지 연산은 동일하게 유지돼요.

동기 (Motivation)

다이어그램에서 보듯 연산이 아주 많고, 각 연산마다 다양한 구현이 가능해요. 이 연산들을 조합해 유효한 fused MoE 구현을 만드는 방법의 수는 금방 감당이 안 될 정도로 불어나요. 모듈러 커널 프레임워크는 연산들을 논리적인 컴포넌트로 묶어서 이 문제를 해결해요. 이렇게 크게 분류해 두면 조합이 관리 가능해지고 코드 중복도 막을 수 있어요. 또 All2All Dispatch&Combine 구현을 fused MoE 구현에서 분리해, 서로 독립적으로 개발·테스트할 수 있게 해줘요. 여기에 더해 모듈러 커널 프레임워크는 각 컴포넌트에 대한 추상 클래스를 제공해서, 앞으로 구현이 따라야 할 뼈대를 분명하게 잡아줘요.

이 문서의 나머지 부분은 Contiguous / Non-Batched 케이스에 초점을 맞춰요. Batched 케이스는 여기서 자연스럽게 확장하면 돼요.

모듈러 커널 컴포넌트

FusedMoEModularKernel은 fused MoE 연산을 3개 부분으로 쪼개요.

  1. TopKWeightAndReduce
  2. FusedMoEPrepareAndFinalizeModular
  3. FusedMoEExpertsModular

TopKWeightAndReduce

TopK 가중치 적용(Weight Application)과 리덕션(Reduction)은 Unpermute 연산 직후, All2All Combine 직전에 일어나요. Unpermute는 FusedMoEExpertsModular가, All2All Combine은 FusedMoEPrepareAndFinalizeModular가 담당한다는 점을 기억해 두세요. TopK 가중치 적용과 리덕션을 FusedMoEExpertsModular 안에서 하는 게 값어치가 있을 때도 있지만, 어떤 구현은 FusedMoEPrepareAndFinalizeModular에서 하기로 하기도 해요. 이런 유연성을 허용하기 위해 TopKWeightAndReduce 추상 클래스를 뒀어요.

TopKWeightAndReduce 구현은 여기 있어요.

FusedMoEPrepareAndFinalizeModular::finalize() 메서드는 TopKWeightAndReduce 인자를 받아서, 그 메서드 안에서 호출해요. FusedMoEModularKernelFusedMoEExpertsModularFusedMoEPrepareAndFinalize 구현 사이의 다리 역할을 하면서, TopK 가중치 적용과 리덕션이 어디서 일어날지를 결정해요.

  • FusedMoEExpertsModular 구현이 가중치 적용과 리덕션을 스스로 한다면, FusedMoEExpertsModular::finalize_weight_and_reduce_impl 메서드는 TopKWeightAndReduceNoOp를 돌려줘요.
  • 그 반대로 FusedMoEExpertsModular 구현이 FusedMoEPrepareAndFinalizeModular::finalize()가 가중치 적용과 리덕션을 해 주길 필요로 한다면, 이 메서드는 TopKWeightAndReduceContiguous / TopKWeightAndReduceNaiveBatched / TopKWeightAndReduceDelegate를 돌려줘요.

FusedMoEPrepareAndFinalizeModular

FusedMoEPrepareAndFinalizeModular 추상 클래스는 prepare, prepare_no_receive, finalize 함수를 노출해요.

  • prepare 함수는 입력 액티베이션 양자화(Quantization)All2All Dispatch를 담당해요.
  • prepare_no_receiveprepare와 비슷하지만, 다른 워커의 결과를 기다리지 않아요. 대신 나중에 호출해서 워커의 최종 결과를 기다리는 "receiver" 콜백을 돌려줘요. 모든 FusedMoEPrepareAndFinalizeModular 클래스가 이 메서드를 지원해야 하는 건 아니지만, 지원한다면 초기 all-to-all 통신과 다른 작업을 인터리브하는 데 쓸 수 있어요(예: shared expert와 fused expert를 섞어 실행).
  • finalize 함수는 All2All Combine을 호출하는 걸 담당해요. 추가로 TopK 가중치 적용과 리덕션을 할 수도 있고 안 할 수도 있어요(TopKWeightAndReduce 섹션 참고).

FusedMoEPrepareAndFinalizeModular Blocks

FusedMoEExpertsModular

FusedMoEExpertsModular 클래스가 MoE 연산의 핵심이 일어나는 곳이에요. 이 추상 클래스는 몇 가지 중요한 함수를 노출해요.

  • apply()
  • workspace_shapes()
  • finalize_weight_and_reduce_impl()

apply()

apply 메서드에서 구현체가 수행하는 작업은 다음과 같아요.

  • Permute
  • 가중치 W1과의 Matmul
  • Act + Mul
  • 양자화(Quantization)
  • 가중치 W2와의 Matmul
  • Unpermute
  • (경우에 따라) TopK 가중치 적용 + 리덕션

workspace_shapes()

핵심 fused MoE 구현은 일련의 연산을 수행해요. 각 연산마다 출력 메모리를 따로 만드는 건 비효율적이에요. 그래서 구현체는 workspace_shapes() 메서드의 출력으로 워크스페이스 shape 2개, 워크스페이스 데이터 타입, fused MoE 출력 shape을 선언해야 해요. 이 정보는 FusedMoEModularKernel::forward()에서 워크스페이스 텐서와 출력 텐서를 할당하는 데 쓰이고, FusedMoEExpertsModular::apply() 메서드에 전달돼요. 그러면 워크스페이스들이 fused MoE 구현의 중간 버퍼로 쓰일 수 있어요.

finalize_weight_and_reduce_impl()

때로는 TopK 가중치 적용과 리덕션을 FusedMoEExpertsModular::apply() 안에서 하는 게 효율적이에요. 여기에서 예시를 볼 수 있어요. 이런 구현을 돕기 위해 TopKWeightAndReduce 추상 클래스를 뒀어요. TopKWeightAndReduce 섹션을 참고하세요.

FusedMoEExpertsModular::finalize_weight_and_reduce_impl()은 구현이 FusedMoEPrepareAndFinalizeModular::finalize()에서 쓰길 원하는 TopKWeightAndReduce 객체를 돌려줘요.

FusedMoEExpertsModular Blocks

FusedMoEModularKernel

FusedMoEModularKernelFusedMoEPrepareAndFinalizeModularFusedMoEExpertsModular 객체로 구성돼요. FusedMoEModularKernel의 의사코드/스케치는 다음과 같아요.

class FusedMoEModularKernel:
    def __init__(self,
                 prepare_finalize: FusedMoEPrepareAndFinalizeModular,
                 fused_experts: FusedMoEExpertsModular):

        self.prepare_finalize = prepare_finalize
        self.fused_experts = fused_experts

    def forward(self, DP_A):

        Aq, A_scale, _, _, _ = self.prepare_finalize.prepare(DP_A, ...)

        workspace13_shape, workspace2_shape, _, _ = self.fused_experts.workspace_shapes(...)

        # allocate workspaces
        workspace_13 = torch.empty(workspace13_shape, ...)
        workspace_2 = torch.empty(workspace2_shape, ...)

        # execute fused_experts
        fe_out = self.fused_experts.apply(Aq, A_scale, workspace13, workspace2, ...)

        # war_impl is an object of type TopKWeightAndReduceNoOp if the fused_experts implementations
        # performs the TopK Weight Application and Reduction.
        war_impl = self.fused_experts.finalize_weight_and_reduce_impl()

        output = self.prepare_finalize.finalize(fe_out, war_impl,...)

        return output

어떻게 하나 (How-To)

FusedMoEPrepareAndFinalizeModular 타입 추가하기

보통 FusedMoEPrepareAndFinalizeModular 타입은 Anll2All Dispatch&Combine 구현/커널에 기반을 둬요. 예를 들어,

  • DeepEPHTPrepareAndFinalize 타입은 DeepEP High-Throughput All2All 커널에 기반을 두고,
  • DeepEPLLPrepareAndFinalize 타입은 DeepEP Low-Latency All2All 커널에 기반을 둬요.

1단계: All2All 매니저 추가

All2All 매니저의 목적은 All2All 커널 구현을 설정하는 거예요. FusedMoEPrepareAndFinalizeModular 구현들은 보통 All2All 매니저에서 커널 구현 "핸들"을 받아와 Dispatch와 Combine 함수를 호출해요. 여기에서 All2All 매니저 구현을 확인하세요.

2단계: FusedMoEPrepareAndFinalizeModular 타입 추가

이 섹션은 FusedMoEPrepareAndFinalizeModular 추상 클래스가 노출하는 각 함수가 무슨 의미인지 설명해요.

FusedMoEPrepareAndFinalizeModular::prepare(): prepare 메서드는 양자화와 All2All Dispatch를 구현해요. 보통 관련 All2All 매니저의 Dispatch 함수를 호출해요.

FusedMoEPrepareAndFinalizeModular::has_prepare_no_receive(): 이 서브클래스가 prepare_no_receive를 구현하는지 여부를 나타내요. 기본값은 False예요.

FusedMoEPrepareAndFinalizeModular::prepare_no_receive(): prepare_no_receive 메서드는 양자화와 All2All Dispatch를 구현해요. 다만 dispatch 연산의 결과를 기다리지 않고, 대신 나중에 호출해서 최종 결과를 기다릴 수 있는 thunk를 돌려줘요. 보통 관련 All2All 매니저의 Dispatch 함수를 호출해요.

FusedMoEPrepareAndFinalizeModular::finalize(): TopK 가중치 적용과 리덕션, 그리고 All2All Combine을 수행할 수 있어요. 보통 관련 All2AllManager의 Combine 함수를 호출해요.

FusedMoEPrepareAndFinalizeModular::activation_format(): prepare 메서드의 출력(즉 All2All dispatch)이 Batched라면 FusedMoEActivationFormat.BatchedExperts를, 그렇지 않으면 FusedMoEActivationFormat.Standard를 돌려줘요.

FusedMoEPrepareAndFinalizeModular::topk_indices_dtype(): TopK ids의 데이터 타입이에요. 일부 All2All 커널은 TopK ids의 데이터 타입에 엄격한 요구사항이 있어요. 이 요구사항은 FusedMoe::select_experts 함수에 전달되어 지켜지게 해요. 엄격한 요구사항이 없다면 None을 돌려줘요.

FusedMoEPrepareAndFinalizeModular::max_num_tokens_per_rank(): 한 번에 All2All Dispatch에 제출할 수 있는 최대 토큰 수예요.

FusedMoEPrepareAndFinalizeModular::num_dispatchers(): dispatching 유닛의 총 개수예요. 이 값이 Dispatch 출력의 크기를 결정해요. Dispatch 출력 형태는 (num_local_experts, max_num_tokens, K)예요. 여기서 max_num_tokens = num_dispatchers() * max_num_tokens_per_rank()예요.

자신의 All2All 구현과 가장 가까운 기존 FusedMoEPrepareAndFinalizeModular 구현을 골라서 참조로 쓰는 걸 권장해요.

FusedMoEExpertsModular 타입 추가하기

FusedMoEExpertsModular는 fused MoE 연산의 핵심을 수행해요. 추상 클래스가 노출하는 각 함수와 그 의미는 다음과 같아요.

FusedMoEExpertsModular::activation_formats(): 지원하는 입력·출력 액티베이션 형태를 돌려줘요. 즉 Contiguous / Batched 형태예요.

FusedMoEExpertsModular::supports_expert_map(): expert map을 지원하면 True를 돌려줘요.

FusedMoEExpertsModular::workspace_shapes() / FusedMoEExpertsModular::finalize_weight_and_reduce_impl / FusedMoEExpertsModular::apply: 위 FusedMoEExpertsModular 섹션을 참고하세요.

FusedMoEModularKernel 초기화

FusedMoEMethodBase 클래스에는 FusedMoEModularKernel 객체를 만드는 데 함께 관여하는 메서드 3개가 있어요.

  • maybe_make_prepare_finalize,
  • select_gemm_impl, 그리고
  • init_prepare_finalize

maybe_make_prepare_finalize

maybe_make_prepare_finalize 메서드는 현재 all2all 백엔드에 따라 적절한 경우 FusedMoEPrepareAndFinalizeModular 인스턴스를 만드는 책임을 져요(예: EP + DP가 켜진 경우). 기본 클래스 메서드는 현재 EP+DP 케이스에 대한 모든 FusedMoEPrepareAndFinalizeModular 객체를 만드는데, 파생 클래스는 이 메서드를 오버라이드해 다른 시나리오의 prepare/finalize 객체를 만들 수 있어요. 예를 들어 ModelOptNvFp4FusedMoE는 EP+TP 케이스에 대해 FlashInferCutlassMoEPrepareAndFinalize를 만들 수 있죠.

구현은 다음에서 확인하세요.

  • ModelOptNvFp4FusedMoE

select_gemm_impl

select_gemm_impl 메서드는 기본 클래스에 정의돼 있지 않아요. 유효하고 적절한 FusedMoEExpertsModular 객체를 만드는 메서드를 구현하는 건 파생 클래스의 몫이에요.

구현은 다음에서 확인하세요.

  • UnquantizedFusedMoEMethod
  • CompressedTensorsW8A8Fp8MoEMethod
  • CompressedTensorsW8A8Fp8MoECutlassMethod
  • Fp8MoEMethod
  • ModelOptNvFp4FusedMoE

파생 클래스들이에요.

init_prepare_finalize

init_prepare_finalize 메서드는 입력과 env 설정에 따라 적절한 FusedMoEPrepareAndFinalizeModular 객체를 만들고, select_gemm_impl에 적절한 FusedMoEExpertsModular 객체를 물어본 뒤 FusedMoEModularKernel 객체를 구성해요.

init_prepare_finalize를 확인하세요.

중요: FusedMoEMethodBase 파생 클래스들은 apply 메서드에서 FusedMoEMethodBase::fused_experts 객체를 사용해요. 설정이 유효한 FusedMoEModularKernel 객체 생성을 허용하면, 그 객체로 FusedMoEMethodBase::fused_experts를 오버라이드해요. 이렇게 하면 파생 클래스들이 어떤 fused MoE 구현이 쓰이는지와 무관해져요.

유닛 테스트 하는 법

FusedMoEModularKernel 유닛 테스트는 test_modular_kernel_combinations.py에 있어요.

유닛 테스트는 FusedMoEPrepareAndFinalizeModularFusedMoEPremuteExpertsUnpermute 타입의 모든 조합을 순회하면서, 호환되는 조합에 대해 정확성 테스트를 실행해요. FusedMoEPrepareAndFinalizeModular / FusedMoEExpertsModular 구현을 추가한다면,

  1. mk_objects.pyMK_ALL_PREPARE_FINALIZE_TYPESMK_FUSED_EXPERT_TYPES에 구현 타입을 각각 추가하세요.
  2. /tests/kernels/moe/modular_kernel_tools/common.pyConfig::is_batched_prepare_finalize(), Config::is_batched_fused_experts(), Config::is_standard_fused_experts(), Config::is_fe_16bit_supported(), Config::is_fe_fp8_supported(), Config::is_fe_block_fp8_supported() 메서드를 갱신하세요.

이렇게 하면 새 구현이 테스트 스위트에 추가돼요.

FusedMoEPrepareAndFinalizeModular & FusedMoEExpertsModular 호환성 확인 방법

유닛 테스트 파일 test_modular_kernel_combinations.py는 독립 실행 스크립트로도 실행할 수 있어요.

예: python3 -m tests.kernels.moe.test_modular_kernel_combinations --pf-type DeepEPLLPrepareAndFinalize --experts-type BatchedTritonExperts

부수 효과로, 이 스크립트로 FusedMoEPrepareAndFinalizeModular & FusedMoEExpertsModular 호환성을 테스트할 수 있어요. 호환되지 않는 타입으로 호출하면 에러가 나요.

프로파일링 하는 법

profile_modular_kernel.py를 확인하세요. 이 스크립트는 호환되는 FusedMoEPrepareAndFinalizeModularFusedMoEExpertsModular 타입 조합에 대해 단일 FusedMoEModularKernel::forward() 호출의 Torch trace를 만들 수 있어요.

예: python3 -m tests.kernels.moe.modular_kernel_tools.profile_modular_kernel --pf-type DeepEPLLPrepareAndFinalize --experts-type BatchedTritonExperts

FusedMoEPrepareAndFinalizeModular 구현

사용 가능한 모든 모듈러 prepare/finalize 서브클래스 목록은 Fused MoE Kernel features를 참고하세요.

FusedMoEExpertsModular

사용 가능한 모든 모듈러 experts 목록은 Fused MoE Kernel features를 참고하세요.

더 알아보기 (Learn more)