Fused MoE Modular Kernel

Fused MoE Modular Kernel

소개 (Introduction)

FusedMoEModularKernel은 여기에 구현돼 있습니다.

입력 활성화(activation) 형식에 따라 fused MoE 구현은 크게 두 유형으로 나뉩니다.

  • Contiguous / Standard / Non-Batched
  • Batched

참고: 문서 전반에서 Contiguous, Standard, Non-Batched는 같은 의미로 사용됩니다.

입력 활성화 형식은 사용 중인 All2All Dispatch에 전적으로 의존합니다.

  • Contiguous 변형에서 All2All Dispatch는 활성화를 형태 (M, K)의 연속 텐서와 형태 (M, num_topk)의 TopK Ids·TopK weights로 반환합니다. 예는 DeepEPHTPrepareAndFinalize를 참고하세요.
  • Batched 변형에서 All2All Dispatch는 활성화를 형태 (num_experts, max_tokens, K)의 텐서로 반환합니다. 같은 전문가에 배정된 활성화/토큰이 함께 배칭됩니다. 텐서의 모든 항목이 유효하진 않습니다. 활성화 텐서는 보통 크기 num_expertsexpert_num_tokens 텐서와 함께 오는데, expert_num_tokens[i]는 i번째 전문가에 배정된 유효 토큰 수를 나타냅니다. 예는 DeepEPLLPrepareAndFinalize를 참고하세요.

fused MoE 연산은 일반적으로 Contiguous와 Batched 변형 모두에서 여러 연산으로 이뤄집니다.

참고: 연산 측면에서 Batched와 Non-Batched의 주요 차이는 Permute/Unpermute 연산입니다. 나머지 연산은 모두 유지됩니다.

출처: 문서

본문

동기 (Motivation)

다이어그램에서 볼 수 있듯 연산이 매우 많고 각 연산마다 다양한 구현이 가능합니다. 유효한 fused MoE 구현을 만들기 위해 연산을 조합하는 방법의 집합은 곧 다루기 어려워집니다. Modular Kernel 프레임워크는 연산을 논리적 컴포넌트로 그룹화해 이 문제를 해결합니다. 이 광범위한 분류는 조합을 관리 가능하게 만들고 코드 중복을 방지합니다. 또한 All2All Dispatch·Combine 구현을 fused MoE 구현에서 분리해 독립적인 개발·테스트를 가능하게 합니다. 나아가 Modular Kernel 프레임워크는 각 컴포넌트에 대해 Abstract 클래스를 도입해 향후 구현을 위한 잘 정의된 골격을 제공합니다.

나머지 문서는 Contiguous / Non-Batched 사례에 초점을 맞춥니다. Batched 사례로의 확장은 단순합니다.

ModularKernel 컴포넌트 (ModularKernel Components)

FusedMoEModularKernel은 fused MoE 연산을 3부분으로 나눕니다:

  • TopKWeightAndReduce
  • FusedMoEPrepareAndFinalizeModular
  • FusedMoEExpertsModular

TopKWeightAndReduce

TopK 가중치 적용·감소(TopK Weight Application and Reduction) 컴포넌트는 Unpermute 연산 직후, All2All Combine 전에 일어납니다. FusedMoEExpertsModular가 Unpermute를 담당하고 FusedMoEPrepareAndFinalizeModular가 All2All Combine을 담당합니다. TopK 가중치 적용·감소를 FusedMoEExpertsModular에서 하는 것도 가치가 있습니다. 그러나 일부 구현은 FusedMoEPrepareAndFinalizeModular에서 하기로 선택합니다. 이 유연성을 가능하게 하기 위해 TopKWeightAndReduce 추상 클래스가 있습니다.

TopKWeightAndReduce 구현은 여기에서 찾을 수 있습니다.

FusedMoEPrepareAndFinalizeModular::finalize() 메서드는 메서드 안에서 호출되는 TopKWeightAndReduce 인자를 받습니다. FusedMoEModularKernelFusedMoEExpertsModularFusedMoEPrepareAndFinalize 구현 사이의 다리 역할을 하며 TopK 가중치 적용·감소가 어디서 일어날지 결정합니다.

FusedMoEPrepareAndFinalizeModular

FusedMoEPrepareAndFinalizeModular 추상 클래스는 prepare, prepare_no_receive, finalize 함수를 노출합니다. prepare 함수는 입력 활성화 양자화와 All2All Dispatch를 담당합니다. 구현된 경우 prepare_no_receive는 다른 워커의 결과를 기다리지 않는 점만 빼고 prepare와 같습니다. 대신 워커의 최종 결과를 기다리기 위해 호출해야 하는 "receiver" 콜백을 반환합니다. 이 메서드가 모든 FusedMoEPrepareAndFinalizeModular 클래스에서 지원돼야 하는 것은 아니지만, 사용 가능하면 초기 all to all 통신과 작업을 인터리브하는 데 쓸 수 있습니다(예: shared experts와 fused experts 인터리빙). finalize 함수는 All2All Combine 호출을 담당합니다. 추가로 finalize 함수는 TopK 가중치 적용·감소를 할 수도 있고 하지 않을 수도 있습니다(TopKWeightAndReduce 절 참고).

FusedMoEExpertsModular

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

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

apply()

apply 메서드는 구현이 수행하는 곳입니다:

  • Permute
  • 가중치 W1과의 Matmul
  • Act + Mul
  • 양자화
  • 가중치 W2와의 Matmul
  • Unpermute
  • 선택적 TopK 가중치 적용 + 감소

workspace_shapes()

핵심 fused MoE 구현은 일련의 연산을 수행합니다. 각 연산에 개별적으로 출력 메모리를 만드는 것은 비효율적입니다. 이를 위해 구현은 workspace_shapes() 메서드의 출력으로 workspace 형태 2개와 workspace 데이터타입, fused MoE 출력 형태를 선언해야 합니다. 이 정보는 FusedMoEModularKernel::forward()에서 workspace 텐서와 출력 텐서를 할당하고 FusedMoEExpertsModular::apply() 메서드에 전달하는 데 사용됩니다. 그런 다음 workspace는 fused MoE 구현에서 중간 버퍼로 사용할 수 있습니다.

finalize_weight_and_reduce_impl()

때로는 FusedMoEExpertsModular::apply() 안에서 TopK 가중치 적용·감소를 수행하는 것이 효율적입니다. 예는 여기를 참고하세요. 이런 구현을 돕기 위해 TopKWeightAndReduce 추상 클래스가 있습니다. TopKWeightAndReduce 절을 참고하세요. FusedMoEExpertsModular::finalize_weight_and_reduce_impl()은 구현이 FusedMoEPrepareAndFinalizeModular::finalize()가 사용하기를 원하는 TopKWeightAndReduce 객체를 반환합니다.

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 타입은 All2All Dispatch & Combine 구현/커널로 뒷받침됩니다. 예:

  • DeepEPHTPrepareAndFinalize 타입은 DeepEP High-Throughput All2All 커널로 뒷받침됩니다.
  • DeepEPLLPrepareAndFinalize 타입은 DeepEP Low-Latency All2All 커널로 뒷받침됩니다.

1단계: All2All 매니저 추가 (Add an All2All manager)

All2All Manager의 목적은 All2All 커널 구현을 설정하는 것입니다. FusedMoEPrepareAndFinalizeModular 구현은 보통 All2All Manager에서 커널 구현 "핸들"을 가져와 Dispatch·Combine 함수를 호출합니다. All2All Manager 구현은 여기를 참고하세요.

2단계: FusedMoEPrepareAndFinalizeModular 타입 추가 (Add a FusedMoEPrepareAndFinalizeModular Type)

이 절은 FusedMoEPrepareAndFinalizeModular 추상 클래스가 노출하는 다양한 함수의 의미를 설명합니다.

FusedMoEPrepareAndFinalizeModular::prepare(): 준비 메서드는 양자화와 All2All Dispatch를 구현합니다. 보통 관련 All2All Manager의 Dispatch 함수가 호출됩니다.

FusedMoEPrepareAndFinalizeModular::has_prepare_no_receive(): 이 서브클래스가 prepare_no_receive를 구현했는지 여부를 나타냅니다. 기본값 False.

FusedMoEPrepareAndFinalizeModular::prepare_no_receive(): 양자화와 All2All Dispatch를 구현합니다. dispatch 연산의 결과를 기다리지 않고, 최종 결과를 기다리기 위해 호출할 수 있는 thunk를 반환합니다. 보통 관련 All2All Manager의 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(): 디스패칭 유닛의 총 수. 이 값은 Dispatch 출력의 크기를 결정합니다. Dispatch 출력의 형태는 (num_local_experts, max_num_tokens, K)입니다. 여기서 max_num_tokens = num_dispatchers() * max_num_tokens_per_rank()입니다.

기존 FusedMoEPrepareAndFinalizeModular 구현 중 All2All 구현과 가장 가까운 것을 참고 자료로 선택하길 권장합니다.

FusedMoEExpertsModular 타입 추가하기 (How To Add a FusedMoEExpertsModular Type)

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 초기화 (FusedMoEModularKernel Initialization)

FusedMoEMethodBase 클래스는 FusedMoEModularKernel 객체를 만드는 일을 공동으로 담당하는 3개 메서드가 있습니다:

  • maybe_make_prepare_finalize
  • select_gemm_impl
  • init_prepare_finalize

maybe_make_prepare_finalize

maybe_make_prepare_finalize 메서드는 현재 all2all 백엔드(예: EP + DP 활성화 시)에 따라 적절할 때 FusedMoEPrepareAndFinalizeModular 인스턴스를 구성합니다. 기본 클래스 메서드는 현재 EP+DP 사례의 모든 FusedMoEPrepareAndFinalizeModular 객체를 구성합니다. 파생 클래스는 이 메서드를 오버라이드해 다른 시나리오의 prepare/finalize 객체를 구성할 수 있습니다. 예: ModelOptNvFp4FusedMoE는 EP+TP 사례의 FlashInferCutlassMoEPrepareAndFinalize를 구성할 수 있습니다. 구현은 다음을 참고하세요:

select_gemm_impl

select_gemm_impl 메서드는 기본 클래스에서 정의되지 않습니다. 파생 클래스가 유효한/적절한 FusedMoEExpertsModular 객체를 구성하는 메서드를 구현하는 것이 책임입니다. 구현은 다음을 참고하세요:

init_prepare_finalize

입력·환경 설정에 따라 init_prepare_finalize 메서드는 적절한 FusedMoEPrepareAndFinalizeModular 객체를 만들고, 그다음 select_gemm_impl에 적절한 FusedMoEExpertsModular 객체를 질의해 FusedMoEModularKernel 객체를 구축합니다.

init_prepare_finalize를 참고하세요. 중요: FusedMoEMethodBase 파생 클래스는 apply 메서드에서 FusedMoEMethodBase::fused_experts 객체를 사용합니다. 설정이 유효한 FusedMoEModularKernel 객체 구성을 허용하면 그것으로 FusedMoEMethodBase::fused_experts를 오버라이드합니다. 이는 파생 클래스를 사용 중인 fused MoE 구현과 무관하게 만듭니다.

단위 테스트 방법 (How To Unit Test)

FusedMoEModularKernel 단위 테스트는 test_modular_kernel_combinations.py에 있습니다.

단위 테스트는 FusedMoEPrepareAndFinalizeModularFusedMoEPremuteExpertsUnpermute 타입의 모든 조합을 순회하며, 호환되면 일부 정확성 테스트를 실행합니다. FusedMoEPrepareAndFinalizeModular / FusedMoEExpertsModular 구현을 추가한다면:

  • 구현 타입을 mk_objects.pyMK_ALL_PREPARE_FINALIZE_TYPESMK_FUSED_EXPERT_TYPES에 각각 추가합니다.
  • /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 호환성 확인 (How To Check Compatibility)

test_modular_kernel_combinations.py 단위 테스트 파일은 독립 실행 스크립트로도 실행할 수 있습니다. 예: python3 -m tests.kernels.moe.test_modular_kernel_combinations --pf-type DeepEPLLPrepareAndFinalize --experts-type BatchedTritonExperts 부수적으로 이 스크립트는 FusedMoEPrepareAndFinalizeModular & FusedMoEExpertsModular 호환성 테스트에 쓸 수 있습니다. 호환되지 않는 타입으로 호출하면 오류가 납니다.

프로파일링 방법 (How To Profile)

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 구현 (Implementations)

사용 가능한 모든 modular prepare/finalize 서브클래스 목록은 Fused MoE Kernel 기능을 참고하세요.

FusedMoEExpertsModular

사용 가능한 모든 modular experts 목록은 Fused MoE Kernel 기능을 참고하세요.

더 알아보기 (Learn more)