diff --git a/docs/design_docs/flashinfer_moe_api.md b/docs/design_docs/flashinfer_moe_api.md index ac1fc4b69eb..c73f7d52699 100644 --- a/docs/design_docs/flashinfer_moe_api.md +++ b/docs/design_docs/flashinfer_moe_api.md @@ -42,18 +42,19 @@ config = MoEConfig( ), quant=QuantConfig(QuantDtype.FP4, QuantGranularity.BlockScale), experts=ExpertConfig(intermediate_size=2048, local_num_experts=32), - backends=[TrtllmFp4Config(extra_backend_params...), CutlassConfig(extra_backend_params...)], + backends=[TrtllmFp4Config(extra_backend_params...), CutlassNvfp4Config()], ) # --- Find possible backends --- backends = MoELayer.find_backends(**config) -# this contains {"trtllm_fp4":TrtllmFp4Config(), "cutlass_fp4":CutlassConfig()} -# or {"trtllm_fp4":"unsupported reason...", "cutlass_fp4":CutlassConfig()} +# this contains {"trtllm_fp4":TrtllmFp4Config(), "cutlass_nvfp4":CutlassNvfp4Config()} +# or {"trtllm_fp4":"unsupported reason...", "cutlass_nvfp4":CutlassNvfp4Config()} # more modification to the backends' parameters could be done here -backends=["trtllm_fp4":TrtllmFp4Config(extra_backend_params...),"cutlass_fp4":CutlassConfig(extra_backend_params...)] +backends=["trtllm_fp4":TrtllmFp4Config(extra_backend_params...),"cutlass_nvfp4":CutlassNvfp4Config()] # --- Prepare Inputs Data --- weight_pack = MoEWeightPack() # the data is possibly obtained through helper functions then added here -weight_pack.prepare_for("trtllm_fp4", trtllm_weights) weight_pack.prepare_for("cutlass_fp4", cutlass_weights) +weight_pack.prepare_for("trtllm_fp4", trtllm_weights) +weight_pack.prepare_for("cutlass_nvfp4", CutlassNvfp4Config.prepare_weights(...)) act_pack = MoEActivationPack( hidden_states_q=cute_dsl_data["x"], hidden_states_scale=x_sf, @@ -120,9 +121,8 @@ Individual backend configs provided in an ordered list. The autotuner or heurist # Single backend backends = [TrtllmFp4Config()] # Multiple candidates — autotuner or heuristic picks best -backends = [TrtllmFp4Config(), TrtllmFp8BlockConfig(), CutlassConfig()] +backends = [TrtllmFp4Config(), CutlassNvfp4Config()] # | is associative, returns BackendOptions -# CutlassConfig is always the universal fallback ``` Each backend config declares its own preconditions: @@ -131,11 +131,11 @@ Each backend config declares its own preconditions: class TrtllmFp4Config: @classmethod def supported(cls, arch: int) -> bool: - return arch >= 90 # Hopper+ -class CutlassConfig: + return arch in (100, 103, 107) +class CutlassNvfp4Config: @classmethod def supported(cls, arch: int) -> bool: - return True # universal fallback + return arch in (100, 103, 107, 110, 120, 121) ``` ### 3.3 MoEConfig — \*\*unpack protocol @@ -147,7 +147,7 @@ config = MoEConfig( routing=RoutingConfig(num_experts=256, top_k=8, method=RoutingMethodType.DeepSeekV3), quant=QuantConfig(QuantDtype.FP4, QuantGranularity.BlockScale), experts=ExpertConfig(intermediate_size=2048, local_num_experts=32), - backends=[TrtllmFp4Config(), CutlassConfig()], + backends=[TrtllmFp4Config(), CutlassNvfp4Config()], ) # Unpack into any call accepting these kwargs output = moe_layer(tensors, **config) @@ -258,7 +258,7 @@ repro.benchmark_all() # FinalizeConfig 0.11ms # Total 3.13ms # Isolate backend to narrow regression -repro.isolate_backend(CutlassConfig()) +repro.isolate_backend(CutlassNvfp4Config()) repro.isolate_backend(TrtllmFp4Config()) ``` @@ -297,7 +297,7 @@ flashinfer/ __init__.py # BACKEND_REGISTRY, DEFAULT_PRIORITY trtllm_fp4.py # TrtllmFp4Config + adapter trtllm_fp8.py # Fp8Block + Fp8PerTensor + adapters - cutlass.py # CutlassConfig + adapter + cutlass.py # CutlassBf16 / W4A16 / Nvfp4 / Fp8 / Mxfp8 / W4A8 / Humming + adapters repro.py # MoERepro tensors.py # MoETensors, Gemm1Tensors, Gemm2Tensors ``` @@ -689,7 +689,8 @@ out = layer(act, weights) # subsequent calls: cached winner dispatch Key mechanisms (and where they live): - **Two packs, two lifetimes.** `MoEWeightPack` holds long-lived, backend-native weight materializations keyed by `backend_key` (`prepare_for` / `get_view`); `MoEActivationPack` carries per-call pre-routed activations. This is the concrete answer to reviewers' "backends need different weight preprocessing" concern (C29–C32): each backend stores its own view, none is hidden from the caller. -- **First-class prep.** `TrtllmFp4Config.prepare_weights(...)` / `CuteDslConfig.prepare_weights(...)` (backed by `flashinfer/fused_moe/prepare.py`) turn canonical bf16 weights into the native views (C6/C7). +- **First-class prep.** `TrtllmFp4Config.prepare_weights(...)` / `CuteDslConfig.prepare_weights(...)` / `CutlassNvfp4Config.prepare_weights(...)` (and the other quant-specific `Cutlass*Config.prepare_weights` helpers, backed by `flashinfer/fused_moe/prepare.py`) turn canonical bf16 weights into the native views (C6/C7). CUTLASS NVFP4 uses swizzled `fp4_quantize` scales, not the TRTLLM shuffle / BlockMajorK path. CUTLASS FP8 / MXFP8 / W4A8 / Humming likewise keep unshuffled or mixed-input layouts distinct from TRTLLM. Each quant mode uses its matching `Cutlass*Config` / `Cutlass*Runner`; there is no quant-neutral CUTLASS fallback. +- **Breaking change — `CutlassConfig` removed.** The deprecated, unregistered `CutlassConfig` placeholder is gone. It was never a runnable `MoELayer` backend (`supported()` always returned false; it was not in `_BACKEND_RUNNERS`). Import, annotate, serialize, or feature-detect a quant-specific type instead (`CutlassBf16Config`, `CutlassNvfp4Config`, `CutlassFp8PerTensorConfig`, `CutlassFp8BlockConfig`, `CutlassMxfp8Config`, `CutlassMxfp8Mxfp4Config`, `CutlassW4A16Config`, `CutlassW4A8Config`, `CutlassHummingConfig`). Historical **Anchor:** / CR1 quotes earlier in this document still mention `CutlassConfig` as review history, not current API. - **Two-stage cross-backend autotune** (`MoELayer._select_winner`, runners' delegation): for each candidate, the `AutoTuner.choose_one` picks the best *within-backend tactic* (each backend tuned in its own native input schema), then `bench_gpu_time` compares the candidates at their winning tactics and the fastest backend is dispatched. A single `choose_one` over both runners is not possible because their input schemas differ — hence the explicit two stages. - **Winner caching is per token-bucket** (`map_to_hybrid_bucket`): reusing one `MoELayer` across token counts re-selects per bucket; `winner_backend` reports the most-recent choice and `reset_winner()` clears the cache. - **Fail-fast scope** (`MoELayer._validate_mvp_scope`): non-NVFP4 quant or non-Swiglu activation raises `NotImplementedError` at construction. diff --git a/flashinfer/fused_moe/__init__.py b/flashinfer/fused_moe/__init__.py index 9540c8cf459..a2ad884dc75 100644 --- a/flashinfer/fused_moe/__init__.py +++ b/flashinfer/fused_moe/__init__.py @@ -21,9 +21,15 @@ B12xW4A16Config, BackendOptions, CuteDslConfig, - CutlassConfig, CutlassBf16Config, + CutlassFp8BlockConfig, + CutlassFp8PerTensorConfig, + CutlassHummingConfig, + CutlassMxfp8Config, + CutlassMxfp8Mxfp4Config, + CutlassNvfp4Config, CutlassW4A16Config, + CutlassW4A8Config, ExecutionConfig, ExpertConfig, MoEActivationPack, @@ -49,7 +55,14 @@ B12xNvfp4Runner, B12xW4A16Runner, CutlassBf16Runner, + CutlassFp8BlockRunner, + CutlassFp8PerTensorRunner, + CutlassHummingRunner, + CutlassMxfp8Mxfp4Runner, + CutlassMxfp8Runner, + CutlassNvfp4Runner, CutlassW4A16Runner, + CutlassW4A8Runner, CuteDslNvfp4Runner, TrtllmBf16RoutedRunner, TrtllmFp4RoutedRunner, @@ -151,15 +164,28 @@ "ActivationConfig", "B12xNvfp4Config", "B12xNvfp4Runner", - "CutlassBf16Runner", - "CutlassW4A16Runner", "B12xW4A16Config", "B12xW4A16Runner", "BackendOptions", "CuteDslConfig", - "CutlassConfig", "CutlassBf16Config", + "CutlassBf16Runner", + "CutlassFp8BlockConfig", + "CutlassFp8BlockRunner", + "CutlassFp8PerTensorConfig", + "CutlassFp8PerTensorRunner", + "CutlassHummingConfig", + "CutlassHummingRunner", + "CutlassMxfp8Config", + "CutlassMxfp8Mxfp4Config", + "CutlassMxfp8Mxfp4Runner", + "CutlassMxfp8Runner", + "CutlassNvfp4Config", + "CutlassNvfp4Runner", "CutlassW4A16Config", + "CutlassW4A16Runner", + "CutlassW4A8Config", + "CutlassW4A8Runner", "ExecutionConfig", "ExpertConfig", "CuteDslNvfp4Runner", diff --git a/flashinfer/fused_moe/api.py b/flashinfer/fused_moe/api.py index 0a9d74b0e67..93dd396ea01 100644 --- a/flashinfer/fused_moe/api.py +++ b/flashinfer/fused_moe/api.py @@ -32,7 +32,6 @@ import torch from torch import Tensor -from typing_extensions import deprecated from ..tllm_enums import ActivationType, RoutingInputMode, RoutingMethodType @@ -71,6 +70,8 @@ class QuantVariant(Enum): MXFP4 = 5 # MXFP4 weights x MXFP8 activations (TRTLLM W4A8) MxInt4 = 6 W4A16 = 7 # backend-specific 4-bit weights x BF16 activations + W4A8 = 8 # INT4 weights x FP8 activations (CUTLASS SM90 packed mixed-input) + Humming = 9 # MXFP4 weights x FP8 activations with Humming pre-MMA fusion def __repr__(self) -> str: return f"{type(self).__name__}.{self.name}" @@ -301,6 +302,29 @@ def __repr__(self) -> str: # W4A16 uses Hopper-specific mixed-input weight and scale layouts. _CUTLASS_W4A16_ARCHS = (90,) +# NVFP4 CUTLASS fused MoE matches the flat-API skip: SM100/SM110/SM12x. +# Major 10/11/12 covers SM100/103/107, SM110, and SM120/121. +_CUTLASS_NVFP4_ARCHS = (100, 103, 107, 110, 120, 121) + +# Per-tensor FP8 follows the dense CUTLASS architecture list. +_CUTLASS_FP8_ARCHS = _CUTLASS_BF16_ARCHS + +# DeepSeek-style 128x128 FP8 block scaling is Hopper-only in the flat API. +_CUTLASS_FP8_BLOCK_ARCHS = (90,) + +# MXFP8 activations x MXFP4 weights matches a *wider* flat-API skip than +# MXFP8 x MXFP8: ``capability[0] not in [10, 11, 12]`` (SM100/103/107, SM110, +# SM120/121). Do not collapse this to ``_CUTLASS_MXFP8_ARCHS``. +_CUTLASS_MXFP8_MXFP4_ARCHS = (100, 103, 107, 110, 120, 121) + +# MXFP8 x MXFP8 follows the narrower flat skip: ``capability[0] not in [10]`` +# (SM10x only, including B300 SM103). Not claimed on SM11x/SM12x. +_CUTLASS_MXFP8_ARCHS = (100, 103, 107) + +# INT4 W4A8 and Humming MXFP4 x FP8 are Hopper mixed-input paths. +_CUTLASS_W4A8_ARCHS = (90,) +_CUTLASS_HUMMING_ARCHS = (90,) + @dataclass(frozen=True) class TrtllmFp4Config: @@ -556,31 +580,9 @@ def __repr__(self) -> str: return "TrtllmMxInt4Config()" -@deprecated( - "CutlassConfig is deprecated and non-runnable; use CutlassBf16Config or " - "CutlassW4A16Config instead." -) -@dataclass(frozen=True) -class CutlassConfig: - """Legacy quantization-neutral CUTLASS configuration placeholder. - - .. deprecated:: - Use :class:`CutlassBf16Config` or :class:`CutlassW4A16Config` instead. - - This type is preserved for source compatibility, but it is intentionally - not registered with :class:`MoELayer` and therefore is not runnable. Select - a concrete tensor contract such as :class:`CutlassBf16Config` or - :class:`CutlassW4A16Config` instead. - """ - - @classmethod - def supported(cls, arch: int) -> bool: - # Compatibility-only placeholder: it has no registered runner and must - # never be surfaced as a dispatch candidate by BackendOptions.valid_for(). - return False - - def __repr__(self) -> str: - return "CutlassConfig()" +# Breaking change: the deprecated, unregistered CutlassConfig placeholder was +# removed. It was never a runnable MoELayer backend (supported() always False). +# Use a quant-specific Cutlass*Config below instead. @dataclass(frozen=True) @@ -670,6 +672,315 @@ def __repr__(self) -> str: return "CutlassW4A16Config()" +@dataclass(frozen=True) +class CutlassNvfp4Config: + """CUTLASS NVFP4 backend for SM100 / SM110 / SM12x. + + Packed precomputed routing with SwiGLU and ``do_finalize=True``. Expert + parallelism and shared experts are not supported. Both ``hidden_size`` + and ``intermediate_size`` must be divisible by 16 (the NVFP4 scale-vector + size). Activations stay BF16; the kernel quantizes them internally. + + This config is not in the default backend search list. TRTLLM NVFP4 uses + a quantized activation pack, so the two contracts cannot share one + ``MoEActivationPack``. Select it explicitly with + ``BackendOptions((CutlassNvfp4Config(),))``. + """ + + @classmethod + def supported(cls, arch: int) -> bool: + return arch in _CUTLASS_NVFP4_ARCHS + + @staticmethod + def prepare_weights( + w1_bf16, + w2_bf16, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device=None, + ): + """Quantize canonical BF16 weights into CUTLASS-swizzled NVFP4. + + Uses the flat CUTLASS scale layout (``fp4_quantize`` with swizzled + scales), not the TRTLLM shuffle / BlockMajorK path. + """ + from .prepare import prepare_cutlass_nvfp4_weights + + return prepare_cutlass_nvfp4_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + + def __repr__(self) -> str: + return "CutlassNvfp4Config()" + + +@dataclass(frozen=True) +class CutlassFp8PerTensorConfig: + """CUTLASS per-tensor FP8 backend. + + Activations are prequantized E4M3 with a scalar dequant scale on + ``MoEActivationPack.hidden_states_scale``. Weights stay unshuffled; this is + not the TRTLLM MajorK view. Packed precomputed routing with SwiGLU and + ``do_finalize=True``. Not in the default backend search list: TRTLLM + per-tensor FP8 folds the activation scale into the weight view, so the two + contracts cannot share one pack without a conversion. + """ + + @classmethod + def supported(cls, arch: int) -> bool: + return arch in _CUTLASS_FP8_ARCHS + + @staticmethod + def prepare_weights( + w1_bf16, + w2_bf16, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device=None, + ): + """Quantize canonical BF16 weights into unshuffled per-tensor FP8.""" + from .prepare import prepare_cutlass_fp8_per_tensor_weights + + return prepare_cutlass_fp8_per_tensor_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + + @staticmethod + def prepare_activations(hidden_states_bf16): + """Quantize BF16 activations to E4M3 plus a scalar dequant scale.""" + from .prepare import prepare_cutlass_fp8_per_tensor_activations + + return prepare_cutlass_fp8_per_tensor_activations(hidden_states_bf16) + + def __repr__(self) -> str: + return "CutlassFp8PerTensorConfig()" + + +@dataclass(frozen=True) +class CutlassFp8BlockConfig: + """CUTLASS DeepSeek-style 128x128 FP8 block-scale backend (SM90). + + Activations stay BF16; the kernel quantizes them internally. This is not + the TRTLLM block-FP8 activation pack. Packed precomputed routing with + SwiGLU and ``do_finalize=True``. + """ + + @classmethod + def supported(cls, arch: int) -> bool: + return arch in _CUTLASS_FP8_BLOCK_ARCHS + + @staticmethod + def prepare_weights( + w1_bf16, + w2_bf16, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device=None, + ): + """Quantize canonical BF16 weights into 128x128 FP8 block scales.""" + from .prepare import prepare_cutlass_fp8_block_weights + + return prepare_cutlass_fp8_block_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + + def __repr__(self) -> str: + return "CutlassFp8BlockConfig()" + + +@dataclass(frozen=True) +class CutlassMxfp8Mxfp4Config: + """CUTLASS MXFP8-activation x MXFP4-weight backend for SM100 / SM110 / SM12x. + + Activations are MXFP8 with a swizzled ``input_sf``. Weights are packed + MXFP4 viewed as int64 at launch. Packed precomputed routing with SwiGLU + and ``do_finalize=True``. Both ``hidden_size`` and ``intermediate_size`` + must be divisible by 128. + """ + + @classmethod + def supported(cls, arch: int) -> bool: + return arch in _CUTLASS_MXFP8_MXFP4_ARCHS + + @staticmethod + def prepare_weights( + w1_bf16, + w2_bf16, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device=None, + ): + """Quantize canonical BF16 weights into CUTLASS MXFP4.""" + from .prepare import prepare_cutlass_mxfp8_mxfp4_weights + + return prepare_cutlass_mxfp8_mxfp4_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + + @staticmethod + def prepare_activations(hidden_states_bf16): + """Quantize BF16 activations to MXFP8 with a swizzled scale buffer.""" + from .prepare import prepare_cutlass_mxfp8_activations + + return prepare_cutlass_mxfp8_activations(hidden_states_bf16) + + def __repr__(self) -> str: + return "CutlassMxfp8Mxfp4Config()" + + +@dataclass(frozen=True) +class CutlassMxfp8Config: + """CUTLASS MXFP8-activation x MXFP8-weight backend (SM100 / SM103 / SM107). + + Activations are MXFP8 with a swizzled ``input_sf``. Weights stay E4M3 with + packed int32 scale tiles. ``hidden_size`` and ``intermediate_size`` must be + divisible by 128 so the gated fc1 scale layout matches the binding. + Packed precomputed routing with SwiGLU and ``do_finalize=True``. + """ + + @classmethod + def supported(cls, arch: int) -> bool: + return arch in _CUTLASS_MXFP8_ARCHS + + @staticmethod + def prepare_weights( + w1_bf16, + w2_bf16, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device=None, + ): + """Quantize canonical BF16 weights into CUTLASS MXFP8.""" + from .prepare import prepare_cutlass_mxfp8_weights + + return prepare_cutlass_mxfp8_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + + @staticmethod + def prepare_activations(hidden_states_bf16): + """Quantize BF16 activations to MXFP8 with a swizzled scale buffer.""" + from .prepare import prepare_cutlass_mxfp8_activations + + return prepare_cutlass_mxfp8_activations(hidden_states_bf16) + + def __repr__(self) -> str: + return "CutlassMxfp8Config()" + + +@dataclass(frozen=True) +class CutlassW4A8Config: + """CUTLASS INT4-weight x FP8-activation backend for SM90. + + Activations stay BF16; the kernel quantizes them internally with the packed + mixed-input INT4 layout. Packed precomputed routing with SwiGLU and + ``do_finalize=True``. + """ + + @classmethod + def supported(cls, arch: int) -> bool: + return arch in _CUTLASS_W4A8_ARCHS + + @staticmethod + def prepare_weights( + w1_bf16, + w2_bf16, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device=None, + ): + """Quantize canonical BF16 weights into SM90 interleaved INT4.""" + from .prepare import prepare_cutlass_w4a8_weights + + return prepare_cutlass_w4a8_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + + def __repr__(self) -> str: + return "CutlassW4A8Config()" + + +@dataclass(frozen=True) +class CutlassHummingConfig: + """CUTLASS Humming MXFP4-weight x FP8-activation backend for SM90. + + Activations stay BF16; weights use Humming pre-MMA E8M0 fusion plus the + SM90 mixed-input interleave. Packed precomputed routing with SwiGLU and + ``do_finalize=True``. + """ + + @classmethod + def supported(cls, arch: int) -> bool: + return arch in _CUTLASS_HUMMING_ARCHS + + @staticmethod + def prepare_weights( + w1_bf16, + w2_bf16, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device=None, + ): + """Quantize canonical BF16 weights into the Humming mixed-input layout.""" + from .prepare import prepare_cutlass_humming_weights + + return prepare_cutlass_humming_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + + def __repr__(self) -> str: + return "CutlassHummingConfig()" + + @dataclass(frozen=True) class CuteDslConfig: """CuteDSL NVFP4 backend — SM100 family only (Blackwell SM100, SM103). @@ -803,9 +1114,15 @@ def __repr__(self) -> str: TrtllmFp8PerTensorConfig, TrtllmBf16Config, TrtllmMxInt4Config, - CutlassConfig, CutlassBf16Config, CutlassW4A16Config, + CutlassNvfp4Config, + CutlassFp8PerTensorConfig, + CutlassFp8BlockConfig, + CutlassMxfp8Mxfp4Config, + CutlassMxfp8Config, + CutlassW4A8Config, + CutlassHummingConfig, CuteDslConfig, B12xNvfp4Config, B12xW4A16Config, @@ -817,9 +1134,15 @@ def __repr__(self) -> str: TrtllmFp8PerTensorConfig, TrtllmBf16Config, TrtllmMxInt4Config, - CutlassConfig, CutlassBf16Config, CutlassW4A16Config, + CutlassNvfp4Config, + CutlassFp8PerTensorConfig, + CutlassFp8BlockConfig, + CutlassMxfp8Mxfp4Config, + CutlassMxfp8Config, + CutlassW4A8Config, + CutlassHummingConfig, CuteDslConfig, B12xNvfp4Config, B12xW4A16Config, diff --git a/flashinfer/fused_moe/core.py b/flashinfer/fused_moe/core.py index 06658fdd050..e75cc9455b8 100644 --- a/flashinfer/fused_moe/core.py +++ b/flashinfer/fused_moe/core.py @@ -1425,9 +1425,7 @@ def cutlass_fused_moe( "FP8 block scaling not yet implemented for Blackwell." ) elif not is_cuda_version_at_least("12.8"): - raise NotImplementedError( - "FP8 block scaling not implemented for CUDA 12.6 or lower." - ) + raise NotImplementedError("FP8 block scaling requires CUDA 12.8 or newer.") if enable_pdl is None: enable_pdl = device_support_pdl(input.device) diff --git a/flashinfer/fused_moe/layer.py b/flashinfer/fused_moe/layer.py index ae9ca6215c0..360ddf84a55 100644 --- a/flashinfer/fused_moe/layer.py +++ b/flashinfer/fused_moe/layer.py @@ -31,7 +31,14 @@ B12xNvfp4Config, B12xW4A16Config, CutlassBf16Config, + CutlassFp8BlockConfig, + CutlassFp8PerTensorConfig, + CutlassHummingConfig, + CutlassMxfp8Config, + CutlassMxfp8Mxfp4Config, + CutlassNvfp4Config, CutlassW4A16Config, + CutlassW4A8Config, CuteDslConfig, MoEActivationPack, MoEConfig, @@ -46,7 +53,14 @@ B12xNvfp4Runner, B12xW4A16Runner, CutlassBf16Runner, + CutlassFp8BlockRunner, + CutlassFp8PerTensorRunner, + CutlassHummingRunner, + CutlassMxfp8Mxfp4Runner, + CutlassMxfp8Runner, + CutlassNvfp4Runner, CutlassW4A16Runner, + CutlassW4A8Runner, CuteDslNvfp4Runner, TrtllmBf16RoutedRunner, TrtllmFp4RoutedRunner, @@ -62,7 +76,14 @@ # typing the list with this Union gives mypy the visibility it needs. _RunnerT = Union[ CutlassBf16Runner, + CutlassFp8BlockRunner, + CutlassFp8PerTensorRunner, + CutlassHummingRunner, + CutlassMxfp8Mxfp4Runner, + CutlassMxfp8Runner, + CutlassNvfp4Runner, CutlassW4A16Runner, + CutlassW4A8Runner, CuteDslNvfp4Runner, TrtllmFp4RoutedRunner, TrtllmBf16RoutedRunner, @@ -76,7 +97,14 @@ # Map backend-config class -> runner class _BACKEND_RUNNERS: Dict[type, Type[_RunnerT]] = { CutlassBf16Config: CutlassBf16Runner, + CutlassFp8BlockConfig: CutlassFp8BlockRunner, + CutlassFp8PerTensorConfig: CutlassFp8PerTensorRunner, + CutlassHummingConfig: CutlassHummingRunner, + CutlassMxfp8Config: CutlassMxfp8Runner, + CutlassMxfp8Mxfp4Config: CutlassMxfp8Mxfp4Runner, + CutlassNvfp4Config: CutlassNvfp4Runner, CutlassW4A16Config: CutlassW4A16Runner, + CutlassW4A8Config: CutlassW4A8Runner, CuteDslConfig: CuteDslNvfp4Runner, TrtllmFp4Config: TrtllmFp4RoutedRunner, TrtllmBf16Config: TrtllmBf16RoutedRunner, diff --git a/flashinfer/fused_moe/prepare.py b/flashinfer/fused_moe/prepare.py index 8bcfe7f7f34..1e79ebe58ad 100644 --- a/flashinfer/fused_moe/prepare.py +++ b/flashinfer/fused_moe/prepare.py @@ -24,8 +24,9 @@ ``TrtllmFp4Config.prepare_weights(...)`` / ``CuteDslConfig.prepare_weights(...)`` / ``TrtllmBf16Config.prepare_weights(...)`` / ... static helpers (see ``api.py``). -The SM90 Humming-style MXFP4 x FP8 helper is currently exposed as a flat helper -for the CUTLASS fused-MoE path. +CUTLASS fused-MoE paths (BF16, NVFP4, per-tensor FP8, DeepSeek block FP8, +MXFP8, W4A16, W4A8, Humming) have matching ``Cutlass*Config.prepare_weights`` +helpers in this module. """ from __future__ import annotations @@ -42,7 +43,7 @@ sm90_mixed_gemm_scale_interleave_trace, sm90_mixed_gemm_weight_interleave_trace, ) -from ..utils import get_compute_capability +from ..utils import get_compute_capability, round_up # Module-level permute-index caches. Permute indices depend on weight geometry # and layout parameters, so matching keys are safe to reuse across calls. @@ -1211,9 +1212,10 @@ def prepare_cutlass_bf16_weights( ) -> Dict[str, torch.Tensor]: """Build the canonical BF16 view consumed by ``CutlassBf16Runner``. - ``w1_bf16`` is ``[E, 2*I, H]`` in semantic ``[up, gate]`` order and - ``w2_bf16`` is ``[E, H, I]``. CUTLASS BF16 paths consume these dense - tensors directly; preparation validates the source contract + ``w1_bf16`` is ``[E, 2*I, H]`` in semantic ``[up, gate]`` order (same as + the flat CUTLASS SwiGLU split that tests name ``(w3, w1)``: up first, gate + second) and ``w2_bf16`` is ``[E, H, I]``. CUTLASS BF16 paths consume these + dense tensors directly; preparation validates the source contract and materializes contiguous tensors on the requested device. """ if device is None: @@ -1344,6 +1346,516 @@ def quantize(weight: torch.Tensor, rows: int, cols: int): } +_NVFP4_SF_VEC_SIZE = 16 +_NVFP4_SF_SWIZZLE_ROWS = 128 + + +def _nvfp4_swizzled_scale_shape(rows: int, cols: int) -> Tuple[int, int]: + """CUTLASS 128x4 swizzled NVFP4 scale layout for an ``[N, K]`` matrix.""" + return ( + round_up(rows, _NVFP4_SF_SWIZZLE_ROWS), + round_up(cols // _NVFP4_SF_VEC_SIZE, 4), + ) + + +def prepare_cutlass_nvfp4_weights( + w1_bf16: torch.Tensor, + w2_bf16: torch.Tensor, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device: Optional[torch.device] = None, +) -> Dict[str, torch.Tensor]: + """Build the CUTLASS NVFP4 view consumed by ``CutlassNvfp4Runner``. + + Canonical source is BF16 ``w1_bf16 [E, 2*I, H]`` in semantic ``[up, gate]`` + order (same as the flat CUTLASS SwiGLU split that tests name ``(w3, w1)``: + up first, gate second) and ``w2_bf16 [E, H, I]``. Each expert is quantized + independently with ``fp4_quantize`` (``sf_vec_size=16``, swizzled scales) + so 128-row swizzle tiles never cross expert boundaries. Global scales are + fixed at 1.0 to match the other unified NVFP4 prepares; the kernel still + receives the six-tensor CUTLASS ``quant_scales`` contract. + + Do not reuse TRTLLM shuffled or BlockMajorK tensors: CUTLASS consumes the + packed uint8 payload plus the swizzled ``fp4_quantize`` scale buffer. + """ + from ..quantization.fp4_quantization import fp4_quantize + + if w1_bf16.dtype != torch.bfloat16 or w2_bf16.dtype != torch.bfloat16: + raise TypeError( + "prepare_cutlass_nvfp4_weights expects BF16 weights, got " + f"w1={w1_bf16.dtype}, w2={w2_bf16.dtype}." + ) + if ( + hidden_size % _NVFP4_SF_VEC_SIZE != 0 + or intermediate_size % _NVFP4_SF_VEC_SIZE != 0 + ): + raise ValueError( + "Cutlass NVFP4 requires hidden_size and intermediate_size " + f"divisible by {_NVFP4_SF_VEC_SIZE}." + ) + expected_w1 = (num_local_experts, 2 * intermediate_size, hidden_size) + expected_w2 = (num_local_experts, hidden_size, intermediate_size) + if tuple(w1_bf16.shape) != expected_w1 or tuple(w2_bf16.shape) != expected_w2: + raise ValueError( + f"weight shapes {tuple(w1_bf16.shape)}/{tuple(w2_bf16.shape)} != " + f"expected {expected_w1}/{expected_w2}." + ) + if device is None: + device = w1_bf16.device + device = torch.device(device) + if device.type != "cuda": + raise ValueError(f"Cutlass NVFP4 preparation requires CUDA, got {device}.") + + w1_bf16 = w1_bf16.to(device).contiguous() + w2_bf16 = w2_bf16.to(device).contiguous() + # Unit global scale keeps prepare aligned with the other unified NVFP4 + # backends. Per-expert amax scales remain a flat-API concern. + global_scale = torch.ones(1, device=device, dtype=torch.float32) + + def quantize_experts(weight: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + packed_rows = [] + scale_rows = [] + scale_shape = _nvfp4_swizzled_scale_shape(weight.shape[1], weight.shape[2]) + for expert in range(num_local_experts): + packed, scale = fp4_quantize( + weight[expert], + global_scale=global_scale, + sf_vec_size=_NVFP4_SF_VEC_SIZE, + sf_use_ue8m0=False, + is_sf_swizzled_layout=True, + ) + packed_rows.append(packed) + scale_rows.append(scale.view(*scale_shape)) + return torch.stack(packed_rows), torch.stack(scale_rows) + + w1_q, w1_block_scale = quantize_experts(w1_bf16) + w2_q, w2_block_scale = quantize_experts(w2_bf16) + ones = torch.ones(num_local_experts, device=device, dtype=torch.float32) + act_global = torch.ones((), device=device, dtype=torch.float32) + return { + "fc1_expert_weights": w1_q, + "fc2_expert_weights": w2_q, + "fc1_act_global_scale": act_global, + "fc1_weight_block_scale": w1_block_scale, + "fc1_dequant_scale": ones, + "fc2_act_global_scale": act_global.clone(), + "fc2_weight_block_scale": w2_block_scale, + "fc2_dequant_scale": ones.clone(), + } + + +def _require_canonical_cutlass_bf16_weights( + w1_bf16: torch.Tensor, + w2_bf16: torch.Tensor, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + name: str, + alignment: Optional[int] = None, + require_cuda: bool = False, + device: Optional[torch.device] = None, +) -> Tuple[torch.Tensor, torch.Tensor, torch.device]: + """Validate canonical BF16 expert weights and move them onto ``device``.""" + if w1_bf16.dtype != torch.bfloat16 or w2_bf16.dtype != torch.bfloat16: + raise TypeError( + f"{name} expects BF16 weights, got w1={w1_bf16.dtype}, w2={w2_bf16.dtype}." + ) + if alignment is not None and ( + hidden_size % alignment != 0 or intermediate_size % alignment != 0 + ): + raise ValueError( + f"{name} requires hidden_size and intermediate_size divisible by " + f"{alignment}." + ) + expected_w1 = (num_local_experts, 2 * intermediate_size, hidden_size) + expected_w2 = (num_local_experts, hidden_size, intermediate_size) + if tuple(w1_bf16.shape) != expected_w1 or tuple(w2_bf16.shape) != expected_w2: + raise ValueError( + f"weight shapes {tuple(w1_bf16.shape)}/{tuple(w2_bf16.shape)} != " + f"expected {expected_w1}/{expected_w2}." + ) + if device is None: + device = w1_bf16.device + device = torch.device(device) + if require_cuda and device.type != "cuda": + raise ValueError(f"{name} requires CUDA, got {device}.") + return w1_bf16.to(device).contiguous(), w2_bf16.to(device).contiguous(), device + + +def prepare_cutlass_fp8_per_tensor_activations( + hidden_states_bf16: torch.Tensor, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Quantize ``[M, H]`` BF16 activations to E4M3 plus a scalar dequant scale.""" + if hidden_states_bf16.dtype != torch.bfloat16 or hidden_states_bf16.dim() != 2: + raise ValueError( + "prepare_cutlass_fp8_per_tensor_activations expects a 2D BF16 tensor, " + f"got shape={tuple(hidden_states_bf16.shape)}, " + f"dtype={hidden_states_bf16.dtype}." + ) + fp8_max = torch.finfo(torch.float8_e4m3fn).max + amax = hidden_states_bf16.float().abs().amax() + dequant = torch.where( + amax > 0, amax / fp8_max, torch.ones_like(amax, dtype=torch.float32) + ).to(torch.float32) + quantized = (hidden_states_bf16.float() / dequant).clamp(-fp8_max, fp8_max) + return quantized.to(torch.float8_e4m3fn), dequant.reshape(()) + + +def prepare_cutlass_fp8_per_tensor_weights( + w1_bf16: torch.Tensor, + w2_bf16: torch.Tensor, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device: Optional[torch.device] = None, +) -> Dict[str, torch.Tensor]: + """Build the unshuffled per-tensor FP8 view for ``CutlassFp8PerTensorRunner``. + + Each expert uses one E4M3 multiplier. The returned ``fc1_dequant`` / + ``fc2_dequant`` tensors are the CUTLASS dequant scales (``amax / fp8_max``), + not TRTLLM's inverted calibration multipliers. + """ + w1_bf16, w2_bf16, device = _require_canonical_cutlass_bf16_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + name="prepare_cutlass_fp8_per_tensor_weights", + device=device, + ) + w1_q, w1_mult = _quantize_fp8_per_expert(w1_bf16) + w2_q, w2_mult = _quantize_fp8_per_expert(w2_bf16) + return { + "fc1_expert_weights": w1_q, + "fc2_expert_weights": w2_q, + "fc1_dequant": (1.0 / w1_mult).contiguous(), + "fc2_dequant": (1.0 / w2_mult).contiguous(), + } + + +def prepare_cutlass_fp8_block_weights( + w1_bf16: torch.Tensor, + w2_bf16: torch.Tensor, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device: Optional[torch.device] = None, +) -> Dict[str, torch.Tensor]: + """Build the DeepSeek 128x128 FP8 block-scale view for CUTLASS. + + Reuses the unified DeepSeek weight quantizer and leaves tensors unshuffled. + ``hidden_size`` and ``intermediate_size`` must be divisible by 128. + """ + w1_bf16, w2_bf16, _device = _require_canonical_cutlass_bf16_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + name="prepare_cutlass_fp8_block_weights", + alignment=128, + device=device, + ) + w1_q, w1_scale = _deepseek_fp8_quantize_weights(w1_bf16) + w2_q, w2_scale = _deepseek_fp8_quantize_weights(w2_bf16) + return { + "fc1_expert_weights": w1_q.contiguous(), + "fc2_expert_weights": w2_q.contiguous(), + "fc1_block_scale": w1_scale.contiguous(), + "fc2_block_scale": w2_scale.contiguous(), + } + + +def prepare_cutlass_mxfp8_activations( + hidden_states_bf16: torch.Tensor, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Quantize ``[M, H]`` BF16 activations to MXFP8 with swizzled scales.""" + from ..quantization.fp8_quantization import mxfp8_quantize + + if hidden_states_bf16.dtype != torch.bfloat16 or hidden_states_bf16.dim() != 2: + raise ValueError( + "prepare_cutlass_mxfp8_activations expects a 2D BF16 tensor, " + f"got shape={tuple(hidden_states_bf16.shape)}, " + f"dtype={hidden_states_bf16.dtype}." + ) + if hidden_states_bf16.shape[1] % 32 != 0: + raise ValueError( + "MXFP8 activations require hidden_size divisible by 32, got " + f"{hidden_states_bf16.shape[1]}." + ) + return mxfp8_quantize(hidden_states_bf16, is_sf_swizzled_layout=True, alignment=32) + + +def _quantize_mxfp4_experts( + weight: torch.Tensor, num_local_experts: int +) -> Tuple[torch.Tensor, torch.Tensor]: + from ..quantization.fp4_quantization import mxfp4_quantize + + packed_rows = [] + scale_rows = [] + for expert in range(num_local_experts): + packed, scale = mxfp4_quantize(weight[expert]) + packed_rows.append(packed) + scale_rows.append(scale) + return torch.stack(packed_rows), torch.stack(scale_rows) + + +def prepare_cutlass_mxfp8_mxfp4_weights( + w1_bf16: torch.Tensor, + w2_bf16: torch.Tensor, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device: Optional[torch.device] = None, +) -> Dict[str, torch.Tensor]: + """Build the CUTLASS MXFP4 weight view consumed with MXFP8 activations. + + MXFP4 block scales are 32-wide, but the fused-MoE binding still requires + ``hidden_size`` and ``intermediate_size`` divisible by 128. + """ + w1_bf16, w2_bf16, device = _require_canonical_cutlass_bf16_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + name="prepare_cutlass_mxfp8_mxfp4_weights", + alignment=128, + require_cuda=True, + device=device, + ) + w1_q, w1_scale = _quantize_mxfp4_experts(w1_bf16, num_local_experts) + w2_q, w2_scale = _quantize_mxfp4_experts(w2_bf16, num_local_experts) + fake_input_scale = torch.ones(num_local_experts, device=device, dtype=torch.float32) + return { + "fc1_expert_weights": w1_q.contiguous(), + "fc2_expert_weights": w2_q.contiguous(), + "fc1_expert_scales": w1_scale.contiguous(), + "fc2_expert_scales": w2_scale.contiguous(), + "fc1_input_scale": fake_input_scale, + "fc2_input_scale": fake_input_scale.clone(), + } + + +def _pack_mxfp8_weight_scales( + scale_u8: torch.Tensor, rows: int, cols: int +) -> torch.Tensor: + # CUTLASS MXFP8 SF N-dim is alignToSfDim(N, 128). Callers must pass the + # already-gated row count (2 * round_up(I, 128) for SwiGLU fc1); combining + # as round_up(2*I, 128) is smaller when I % 128 != 0. + num_experts = scale_u8.size(0) + aligned_rows = round_up(rows, 128) + aligned_k_scales = round_up(cols // 32, 4) + return ( + scale_u8.contiguous() + .view(num_experts, aligned_rows, aligned_k_scales) + .view(torch.int32) + .contiguous() + ) + + +def _quantize_mxfp8_experts( + weight: torch.Tensor, num_local_experts: int +) -> Tuple[torch.Tensor, torch.Tensor]: + from ..quantization.fp8_quantization import mxfp8_quantize + + packed_rows = [] + scale_rows = [] + for expert in range(num_local_experts): + packed, scale = mxfp8_quantize( + weight[expert], is_sf_swizzled_layout=True, alignment=32 + ) + packed_rows.append(packed) + scale_rows.append(scale) + return torch.stack(packed_rows), torch.stack(scale_rows) + + +def prepare_cutlass_mxfp8_weights( + w1_bf16: torch.Tensor, + w2_bf16: torch.Tensor, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device: Optional[torch.device] = None, +) -> Dict[str, torch.Tensor]: + """Build the CUTLASS MXFP8 weight view consumed with MXFP8 activations. + + MXFP8 block scales are 32-wide, but the fused-MoE binding requires + ``hidden_size`` and ``intermediate_size`` divisible by 128. The gated fc1 + SF N-dim is ``2 * round_up(I, 128)``, which only matches + ``mxfp8_quantize``'s ``round_up(2*I, 128)`` output when ``I % 128 == 0``. + """ + w1_bf16, w2_bf16, device = _require_canonical_cutlass_bf16_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + name="prepare_cutlass_mxfp8_weights", + alignment=128, + require_cuda=True, + device=device, + ) + w1_q, w1_scale = _quantize_mxfp8_experts(w1_bf16, num_local_experts) + w2_q, w2_scale = _quantize_mxfp8_experts(w2_bf16, num_local_experts) + fake_input_scale = torch.ones(num_local_experts, device=device, dtype=torch.float32) + return { + "fc1_expert_weights": w1_q.contiguous(), + "fc2_expert_weights": w2_q.contiguous(), + "fc1_expert_scales": _pack_mxfp8_weight_scales( + w1_scale, 2 * round_up(intermediate_size, 128), hidden_size + ), + "fc2_expert_scales": _pack_mxfp8_weight_scales( + w2_scale, hidden_size, intermediate_size + ), + "fc1_input_scale": fake_input_scale, + "fc2_input_scale": fake_input_scale.clone(), + } + + +def _quantize_int4_grouped( + weight: torch.Tensor, group_size: int = 128 +) -> Tuple[torch.Tensor, torch.Tensor]: + """Symmetric INT4 pack with per-group BF16 scales along the last dimension.""" + if weight.ndim != 3 or weight.shape[-1] % group_size != 0: + raise ValueError( + "INT4 grouped quantization requires a 3D tensor whose last dim is " + f"divisible by {group_size}; got {tuple(weight.shape)}." + ) + experts, rows, cols = weight.shape + blocks = weight.to(torch.float32).reshape( + experts, rows, cols // group_size, group_size + ) + amax = blocks.abs().amax(dim=-1) + scale = torch.where(amax > 0, amax / 7.0, torch.ones_like(amax)) + quantized = (blocks / scale.unsqueeze(-1)).round().clamp(-8, 7).to(torch.int8) + quantized = quantized.reshape(experts, rows, cols) + even = (quantized[..., 0::2] & 0xF).to(torch.uint8) + odd = (quantized[..., 1::2] & 0xF).to(torch.uint8) + packed = even | (odd << 4) + return packed.contiguous(), scale.to(torch.bfloat16).contiguous() + + +def prepare_cutlass_w4a8_weights( + w1_bf16: torch.Tensor, + w2_bf16: torch.Tensor, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device: Optional[torch.device] = None, +) -> Dict[str, torch.Tensor]: + """Build the SM90 packed INT4 view for ``CutlassW4A8Runner``. + + Canonical BF16 weights are quantized in groups of 128, then folded with the + mixed-input INT4 interleave. Activation prequant scales are identity so the + kernel can quantize BF16 activations internally. + """ + group_size = 128 + w1_bf16, w2_bf16, device = _require_canonical_cutlass_bf16_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + name="prepare_cutlass_w4a8_weights", + alignment=group_size, + require_cuda=True, + device=device, + ) + w1_packed, w1_scale = _quantize_int4_grouped(w1_bf16, group_size) + w2_packed, w2_scale = _quantize_int4_grouped(w2_bf16, group_size) + w1_il = interleave_moe_weights_for_sm90_mixed_gemm(w1_packed, "int4") + w2_il = interleave_moe_weights_for_sm90_mixed_gemm(w2_packed, "int4") + w1_scale_il = interleave_moe_scales_for_sm90_mixed_gemm(w1_scale, group_size) + w2_scale_il = interleave_moe_scales_for_sm90_mixed_gemm(w2_scale, group_size) + ones_h = torch.ones(hidden_size, device=device, dtype=torch.bfloat16) + ones_i = torch.ones(intermediate_size, device=device, dtype=torch.bfloat16) + ones_e = torch.ones(num_local_experts, device=device, dtype=torch.float32) + empty = torch.empty(0, device=device, dtype=torch.bfloat16) + return { + "fc1_expert_weights": w1_il, + "fc2_expert_weights": w2_il, + "fc1_expert_scales": w1_scale_il, + "fc2_expert_scales": w2_scale_il, + "fc1_act_scale": ones_h, + "fc2_act_scale": ones_i, + "fc1_zero": empty, + "fc2_zero": empty.clone(), + "fc1_alpha": ones_e, + "fc2_alpha": ones_e.clone(), + } + + +def prepare_cutlass_humming_weights( + w1_bf16: torch.Tensor, + w2_bf16: torch.Tensor, + *, + num_local_experts: int, + hidden_size: int, + intermediate_size: int, + device: Optional[torch.device] = None, +) -> Dict[str, torch.Tensor]: + """Build the Humming MXFP4 x FP8 mixed-input view for SM90 CUTLASS. + + Logical MXFP4 is produced with the architecture-independent Torch + quantizer, then rewritten and interleaved by + :func:`preprocess_moe_weights_for_sm90_mixed_gemm_humming`. + """ + humming_epilogue_compensation = 64.0 + w1_bf16, w2_bf16, device = _require_canonical_cutlass_bf16_weights( + w1_bf16, + w2_bf16, + num_local_experts=num_local_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + name="prepare_cutlass_humming_weights", + alignment=128, + require_cuda=True, + device=device, + ) + + def quantize(weight: torch.Tensor, rows: int, cols: int): + packed, scales = _quantize_mxfp4_linear( + weight.view(num_local_experts * rows, cols) + ) + return ( + packed.view(num_local_experts, rows, cols // 2), + scales.view(num_local_experts, rows, cols // 32), + ) + + w1_packed, w1_scale = quantize(w1_bf16, 2 * intermediate_size, hidden_size) + w2_packed, w2_scale = quantize(w2_bf16, hidden_size, intermediate_size) + w1_il, w1_scale_il, w1_residual = ( + preprocess_moe_weights_for_sm90_mixed_gemm_humming(w1_packed, w1_scale) + ) + w2_il, w2_scale_il, w2_residual = ( + preprocess_moe_weights_for_sm90_mixed_gemm_humming(w2_packed, w2_scale) + ) + reserved = torch.ones((), device=device, dtype=torch.float32) + return { + "fc1_expert_weights": w1_il, + "fc2_expert_weights": w2_il, + "fc1_expert_scales": w1_scale_il, + "fc2_expert_scales": w2_scale_il, + "fc1_residual_scale": ( + w1_residual * humming_epilogue_compensation + ).contiguous(), + "fc2_residual_scale": ( + w2_residual * humming_epilogue_compensation + ).contiguous(), + "fc2_act_global": reserved, + } + + def _interleave_linear_and_gate( x: torch.Tensor, group_size: int = 64, dim: int = -1 ) -> torch.Tensor: diff --git a/flashinfer/fused_moe/runners.py b/flashinfer/fused_moe/runners.py index 87614fdb7ed..80e9f3fe476 100644 --- a/flashinfer/fused_moe/runners.py +++ b/flashinfer/fused_moe/runners.py @@ -27,11 +27,24 @@ import torch -from ..autotuner import AutoTuner, DynamicTensorSpec, TunableRunner, TuningConfig -from ..utils import next_positive_power_of_2 +from ..autotuner import ( + AutoTuner, + ConstraintSpec, + DynamicTensorSpec, + TunableRunner, + TuningConfig, +) +from ..utils import next_positive_power_of_2, round_up from .api import ( _CUTLASS_BF16_ARCHS, + _CUTLASS_FP8_ARCHS, + _CUTLASS_FP8_BLOCK_ARCHS, + _CUTLASS_HUMMING_ARCHS, + _CUTLASS_MXFP8_ARCHS, + _CUTLASS_MXFP8_MXFP4_ARCHS, + _CUTLASS_NVFP4_ARCHS, _CUTLASS_W4A16_ARCHS, + _CUTLASS_W4A8_ARCHS, ActivationType, MoEActivationPack, MoEConfig, @@ -337,6 +350,29 @@ def get_cache_key_extras(self, inputs: List[torch.Tensor]) -> tuple: return self._cache_key_extras() +def _mxfp8_swizzled_act_sf_numel(num_tokens: int, hidden_size: int) -> int: + """Byte count of a 128x4-swizzled MXFP8 activation scale buffer.""" + return round_up(num_tokens, 128) * round_up(hidden_size // 32, 4) + + +def _infer_mxfp8_swizzled_act_sf_numel(shapes: list) -> int: + """ConstraintSpec callback: ``input_sf`` numel from hidden_states ``[M, H]``.""" + return _mxfp8_swizzled_act_sf_numel(shapes[1][0], shapes[1][1]) + + +def _require_cutlass_tensor( + tensor: torch.Tensor, + *, + name: str, + dtype: torch.dtype, + shape: tuple[int, ...], +) -> None: + if tensor.dtype is not dtype: + raise TypeError(f"{name} must be {dtype}, got {tensor.dtype}.") + if tuple(tensor.shape) != shape: + raise ValueError(f"{name} shape {tuple(tensor.shape)} != expected {shape}.") + + # --------------------------------------------------------------------------- # CUTLASS runners — dense BF16 and mixed-input W4A16 # --------------------------------------------------------------------------- @@ -356,8 +392,13 @@ class _CutlassRunnerBase(MoERunner): supported_routing_modes = (RoutingInputMode.PackedPrecomputed,) supports_expert_parallelism = False _supported_archs: ClassVar[tuple[int, ...]] + _x_dtype: ClassVar[torch.dtype] = torch.bfloat16 _weight_dtype: ClassVar[torch.dtype] - _use_w4_group_scaling: ClassVar[bool] + _use_w4_group_scaling: ClassVar[bool] = False + _use_deepseek_fp8_block_scale: ClassVar[bool] = False + _use_mxfp8_act_scaling: ClassVar[bool] = False + _use_packed_weights: ClassVar[bool] = False + _use_wfp4afp8_humming: ClassVar[bool] = False _required_weight_keys: ClassVar[tuple[str, ...]] _expected_num_inputs: ClassVar[int] # Keep the best N tactics per GEMM stage, then return their Cartesian @@ -381,6 +422,20 @@ def _check_support(self) -> None: f"SM{self._device_arch}; supported architectures are " f"{self._supported_archs}." ) + if self._use_mxfp8_act_scaling and ( + self.config.quant.swizzled_scale_factors is False + ): + raise NotImplementedError( + f"{type(self).__name__} requires swizzled MXFP8 input_sf; " + "linear scales (swizzled_scale_factors=False) are not supported." + ) + if self._use_deepseek_fp8_block_scale: + from ..jit.cpp_ext import is_cuda_version_at_least + + if not is_cuda_version_at_least("12.8"): + raise NotImplementedError( + "FP8 block scaling requires CUDA 12.8 or newer." + ) def __init__(self, config: MoEConfig, device: torch.device): super().__init__() @@ -419,7 +474,7 @@ def _build(self) -> None: with torch.cuda.device(self.device): module = get_cutlass_fused_moe_module(str(self._device_arch)) self._inner = module.MoERunner( - x_dtype=torch.bfloat16, + x_dtype=self._x_dtype, weight_dtype=self._weight_dtype, output_dtype=torch.bfloat16, top_k=self.config.routing.top_k, @@ -430,15 +485,15 @@ def _build(self) -> None: cluster_size=1, cluster_rank=0, enable_alltoall=False, - use_deepseek_fp8_block_scale=False, + use_deepseek_fp8_block_scale=self._use_deepseek_fp8_block_scale, use_w4_group_scaling=self._use_w4_group_scaling, - use_mxfp8_act_scaling=False, + use_mxfp8_act_scaling=self._use_mxfp8_act_scaling, min_latency_mode=False, enable_pdl=self._enable_pdl, activation_type=self.config.activation.type, - use_packed_weights=False, + use_packed_weights=self._use_packed_weights, use_fused_finalize=self._use_fused_finalize, - use_wfp4afp8_humming=False, + use_wfp4afp8_humming=self._use_wfp4afp8_humming, ) def _prepare_tuning_inputs(self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: @@ -455,6 +510,18 @@ def _prepare_tuning_inputs(self, inputs: List[torch.Tensor]) -> List[torch.Tenso ).unsqueeze(0) inputs[2].copy_((token_offsets * top_k + slots) % num_experts) inputs[3].fill_(1.0 / top_k) + if self._use_mxfp8_act_scaling: + hidden_states = inputs[1] + inputs[-1] = torch.full( + ( + _mxfp8_swizzled_act_sf_numel( + hidden_states.shape[0], hidden_states.shape[1] + ), + ), + 127, + dtype=torch.uint8, + device=hidden_states.device, + ) return inputs def get_valid_tactics(self, inputs: List[torch.Tensor], _profile: Any) -> List[Any]: @@ -528,12 +595,16 @@ def _ensure_workspace(self, num_tokens: int, hidden_size: int) -> None: self.config.experts.intermediate_size, self.config.routing.num_experts, self.config.routing.top_k, - x_dtype=torch.bfloat16, + x_dtype=self._x_dtype, weight_dtype=self._weight_dtype, output_dtype=torch.bfloat16, activation_type=self.config.activation.type, + use_deepseek_fp8_block_scale=self._use_deepseek_fp8_block_scale, use_w4_group_scaling=self._use_w4_group_scaling, + use_mxfp8_act_scaling=self._use_mxfp8_act_scaling, use_fused_finalize=self._use_fused_finalize, + use_packed_weights=self._use_packed_weights, + use_wfp4afp8_humming=self._use_wfp4afp8_humming, device=self.device, ) workspace = torch.empty(size, dtype=torch.uint8, device=self.device) @@ -551,19 +622,16 @@ def pack_inputs( f"{type(self).__name__} supports only PackedPrecomputed routing." ) hidden_states = act.hidden_states_q - if hidden_states.ndim != 2 or hidden_states.dtype is not torch.bfloat16: + if hidden_states.ndim != 2 or hidden_states.dtype is not self._x_dtype: raise TypeError( - f"{type(self).__name__} requires 2D BF16 hidden_states_q, got " - f"shape={tuple(hidden_states.shape)}, dtype={hidden_states.dtype}." + f"{type(self).__name__} requires 2D {self._x_dtype} hidden_states_q, " + f"got shape={tuple(hidden_states.shape)}, dtype={hidden_states.dtype}." ) if hidden_states.device != self.device: raise ValueError( f"hidden_states_q is on {hidden_states.device}, expected {self.device}." ) - if act.hidden_states_scale is not None: - raise ValueError( - f"{type(self).__name__} activations do not use hidden_states_scale." - ) + self._validate_activation_scale(act) num_tokens, hidden_size = hidden_states.shape ceiling = self.config.execution.tune_max_num_tokens @@ -588,34 +656,104 @@ def pack_inputs( f"{self.backend_key} prepared weights are missing {missing}." ) weight_inputs = self._pack_weight_inputs(view, hidden_size) + scale_inputs = self._pack_activation_scale_inputs(act) + + # Token-dynamic dims are only the packed prerouted buffers (output, + # hidden, topk_ids, topk_weights). Per-tensor FP8 dequant is 0-dim; + # MXFP8 input_sf is a swizzled 1-D buffer resized by ConstraintSpec. + # Sniffing shape[0] == num_tokens treated a (1,) scale at M=1 as + # token-dynamic and let autotune replace it with a bucket-sized tensor. + input_idxs: tuple[int, ...] = (0, 1, 2, 3) + dim_idxs: tuple[int, ...] = (0, 0, 0, 0) bucket = map_to_hybrid_bucket( num_tokens, self.config.execution.tune_max_num_tokens ) + constraint_specs: tuple[ConstraintSpec, ...] = () + if self._use_mxfp8_act_scaling: + constraint_specs = ( + ConstraintSpec( + 4 + len(weight_inputs), + 0, + _infer_mxfp8_swizzled_act_sf_numel, + ), + ) self.tuning_config = TuningConfig( dynamic_tensor_specs=( DynamicTensorSpec( - input_idx=(0, 1, 2, 3), - dim_idx=(0, 0, 0, 0), + input_idx=input_idxs, + dim_idx=dim_idxs, gen_tuning_buckets=(bucket,), map_to_tuning_buckets=make_hybrid_bucket_mapper( self.config.execution.tune_max_num_tokens ), ), ), + constraint_specs=constraint_specs, use_cuda_graph=True, inputs_pre_hook=self._prepare_tuning_inputs, ) self._ensure_workspace(bucket, hidden_size) - output = hidden_states.new_empty((num_tokens, hidden_size)) + output = torch.empty( + (num_tokens, hidden_size), + dtype=torch.bfloat16, + device=hidden_states.device, + ) return [ output, hidden_states, act.topk_ids, act.topk_weights, *weight_inputs, + *scale_inputs, ] + def _validate_activation_scale(self, act: MoEActivationPack) -> None: + if self._use_mxfp8_act_scaling: + scale = act.hidden_states_scale + num_tokens, hidden_size = act.hidden_states_q.shape + expected = _mxfp8_swizzled_act_sf_numel(num_tokens, hidden_size) + if ( + scale is None + or scale.dtype is not torch.uint8 + or scale.numel() != expected + or not scale.is_contiguous() + ): + got = ( + None + if scale is None + else (scale.dtype, tuple(scale.shape), scale.is_contiguous()) + ) + raise ValueError( + f"{type(self).__name__} requires a contiguous uint8 swizzled " + f"input_sf with {expected} elements for M={num_tokens}, " + f"H={hidden_size}; got {got}." + ) + return + if self._x_dtype is torch.float8_e4m3fn: + scale = act.hidden_states_scale + if scale is None or scale.dim() != 0 or scale.dtype is not torch.float32: + raise ValueError( + f"{type(self).__name__} requires a 0-dim float32 " + "hidden_states_scale dequant factor." + ) + return + if act.hidden_states_scale is not None: + raise ValueError( + f"{type(self).__name__} activations do not use hidden_states_scale." + ) + + def _pack_activation_scale_inputs( + self, act: MoEActivationPack + ) -> List[torch.Tensor]: + if self._use_mxfp8_act_scaling or self._x_dtype is torch.float8_e4m3fn: + assert act.hidden_states_scale is not None + scale = act.hidden_states_scale + if self._use_mxfp8_act_scaling: + scale = scale.reshape(-1) + return [scale] + return [] + def _pack_weight_inputs( self, view: dict[str, torch.Tensor], hidden_size: int ) -> List[torch.Tensor]: @@ -624,6 +762,11 @@ def _pack_weight_inputs( def _quant_scales(self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: return [] + def _input_sf(self, inputs: List[torch.Tensor]) -> torch.Tensor | None: + if self._use_mxfp8_act_scaling: + return inputs[-1] + return None + def _validate_weight_storage(self, tensors: tuple[torch.Tensor, ...]) -> None: if any(t.device != self.device for t in tensors): raise ValueError("CUTLASS prepared weights must match the runner device.") @@ -695,12 +838,18 @@ def forward( inputs[5], output_dtype=torch.bfloat16, quant_scales=self._quant_scales(inputs), + input_sf=self._input_sf(inputs), output=inputs[0], tune_max_num_tokens=self.config.execution.tune_max_num_tokens, enable_pdl=self._enable_pdl, activation_type=self.config.activation.type, + use_deepseek_fp8_block_scale=self._use_deepseek_fp8_block_scale, use_w4_group_scaling=self._use_w4_group_scaling, + use_mxfp8_act_scaling=self._use_mxfp8_act_scaling, + use_packed_weights=self._use_packed_weights, + use_wfp4afp8_humming=self._use_wfp4afp8_humming, use_fused_finalize=self._use_fused_finalize, + swizzled_input_sf=True, profile_ids=profile_ids, workspace_buffer=self._workspace, ) @@ -804,6 +953,583 @@ def _quant_scales(self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: return [inputs[6].view(torch.int32), inputs[7].view(torch.int32)] +class CutlassNvfp4Runner(_CutlassRunnerBase): + """Unified adapter for CUTLASS NVFP4 fused MoE. + + Weights stay packed uint8 in the ``MoEWeightPack`` view. At launch they are + viewed as ``int64``, matching the flat ``cutlass_fused_moe`` NVFP4 ABI; the + inner ``MoERunner`` selects the NVFP4 kernel from that dtype. Activations + remain BF16 and are quantized inside the kernel with unit global scale. + """ + + backend_key = "cutlass_nvfp4" + supported_quant_variants = (QuantVariant.NVFP4,) + _supported_archs = _CUTLASS_NVFP4_ARCHS + _weight_dtype = torch.int64 + _use_w4_group_scaling = False + _required_weight_keys = ( + "fc1_expert_weights", + "fc2_expert_weights", + "fc1_act_global_scale", + "fc1_weight_block_scale", + "fc1_dequant_scale", + "fc2_act_global_scale", + "fc2_weight_block_scale", + "fc2_dequant_scale", + ) + _expected_num_inputs = 12 + + def _pack_weight_inputs( + self, view: dict[str, torch.Tensor], hidden_size: int + ) -> List[torch.Tensor]: + ( + w1, + w2, + a1_gs, + w1_scale, + fc1_dequant, + a2_gs, + w2_scale, + fc2_dequant, + ) = (view[key] for key in self._required_weight_keys) + num_experts = self.config.routing.num_experts + intermediate_size = self.config.experts.intermediate_size + if hidden_size % 16 != 0 or intermediate_size % 16 != 0: + raise ValueError( + "Cutlass NVFP4 requires hidden_size and intermediate_size " + f"divisible by 16, got H={hidden_size}, I={intermediate_size}." + ) + expected_w1 = (num_experts, 2 * intermediate_size, hidden_size // 2) + expected_w2 = (num_experts, hidden_size, intermediate_size // 2) + expected_s1 = ( + num_experts, + round_up(2 * intermediate_size, 128), + round_up(hidden_size // 16, 4), + ) + expected_s2 = ( + num_experts, + round_up(hidden_size, 128), + round_up(intermediate_size // 16, 4), + ) + if w1.dtype is not torch.uint8 or w2.dtype is not torch.uint8: + raise TypeError("Cutlass NVFP4 packed weights must be uint8.") + if w1_scale.dtype is not torch.uint8 or w2_scale.dtype is not torch.uint8: + raise TypeError("Cutlass NVFP4 block scales must be uint8.") + if any( + t.dtype is not torch.float32 + for t in (a1_gs, a2_gs, fc1_dequant, fc2_dequant) + ): + raise TypeError("Cutlass NVFP4 global and dequant scales must be float32.") + if (tuple(w1.shape), tuple(w2.shape)) != (expected_w1, expected_w2): + raise ValueError( + f"Cutlass NVFP4 weight shapes {tuple(w1.shape)}/{tuple(w2.shape)} " + f"!= expected {expected_w1}/{expected_w2}." + ) + if (tuple(w1_scale.shape), tuple(w2_scale.shape)) != ( + expected_s1, + expected_s2, + ): + raise ValueError( + "Cutlass NVFP4 scale shapes " + f"{tuple(w1_scale.shape)}/{tuple(w2_scale.shape)} != expected " + f"{expected_s1}/{expected_s2}." + ) + expected_dequant = (num_experts,) + if ( + tuple(fc1_dequant.shape) != expected_dequant + or tuple(fc2_dequant.shape) != expected_dequant + ): + raise ValueError( + "Cutlass NVFP4 dequant scale shapes " + f"{tuple(fc1_dequant.shape)}/{tuple(fc2_dequant.shape)} != " + f"expected {expected_dequant}." + ) + self._validate_weight_storage( + (w1, w2, a1_gs, w1_scale, fc1_dequant, a2_gs, w2_scale, fc2_dequant) + ) + # Flat NVFP4 CUTLASS selects the kernel from weight dtype int64. + return [ + w1.view(torch.int64), + w2.view(torch.int64), + a1_gs, + w1_scale, + fc1_dequant, + a2_gs, + w2_scale, + fc2_dequant, + ] + + def _quant_scales(self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: + return [ + inputs[6], + inputs[7].view(torch.int32), + inputs[8], + inputs[9], + inputs[10].view(torch.int32), + inputs[11], + ] + + +class CutlassFp8PerTensorRunner(_CutlassRunnerBase): + """Unified adapter for CUTLASS per-tensor FP8 fused MoE.""" + + backend_key = "cutlass_fp8_per_tensor" + supported_quant_variants = (QuantVariant.FP8PerTensor,) + _supported_archs = _CUTLASS_FP8_ARCHS + _x_dtype = torch.float8_e4m3fn + _weight_dtype = torch.float8_e4m3fn + _required_weight_keys = ( + "fc1_expert_weights", + "fc2_expert_weights", + "fc1_dequant", + "fc2_dequant", + ) + _expected_num_inputs = 9 + + def _pack_weight_inputs( + self, view: dict[str, torch.Tensor], hidden_size: int + ) -> List[torch.Tensor]: + w1, w2, w1_dequant, w2_dequant = ( + view[key] for key in self._required_weight_keys + ) + num_experts = self.config.routing.num_experts + intermediate_size = self.config.experts.intermediate_size + expected_w1 = (num_experts, 2 * intermediate_size, hidden_size) + expected_w2 = (num_experts, hidden_size, intermediate_size) + if w1.dtype is not torch.float8_e4m3fn or w2.dtype is not torch.float8_e4m3fn: + raise TypeError("Cutlass FP8 prepared weights must be float8_e4m3fn.") + if tuple(w1.shape) != expected_w1 or tuple(w2.shape) != expected_w2: + raise ValueError( + f"Cutlass FP8 weight shapes {tuple(w1.shape)}/{tuple(w2.shape)} " + f"!= expected {expected_w1}/{expected_w2}." + ) + if ( + w1_dequant.dtype is not torch.float32 + or w2_dequant.dtype is not torch.float32 + ): + raise TypeError("Cutlass FP8 dequant scales must be float32.") + expected_scale = (num_experts,) + if ( + tuple(w1_dequant.shape) != expected_scale + or tuple(w2_dequant.shape) != expected_scale + ): + raise ValueError( + "Cutlass FP8 dequant scale shapes " + f"{tuple(w1_dequant.shape)}/{tuple(w2_dequant.shape)} != " + f"expected {expected_scale}." + ) + self._validate_weight_storage((w1, w2, w1_dequant, w2_dequant)) + return [w1, w2, w1_dequant, w2_dequant] + + def _quant_scales(self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: + act_scale = inputs[8] + gemm2_act_quant = torch.ones((), device=act_scale.device, dtype=torch.float32) + return [ + (inputs[6] * act_scale).float(), + gemm2_act_quant, + inputs[7].float(), + act_scale, + ] + + +class CutlassFp8BlockRunner(_CutlassRunnerBase): + """Unified adapter for CUTLASS DeepSeek 128x128 FP8 block-scale MoE.""" + + backend_key = "cutlass_fp8_block" + supported_quant_variants = (QuantVariant.DeepSeekFp8,) + _supported_archs = _CUTLASS_FP8_BLOCK_ARCHS + _weight_dtype = torch.float8_e4m3fn + _use_deepseek_fp8_block_scale = True + _required_weight_keys = ( + "fc1_expert_weights", + "fc2_expert_weights", + "fc1_block_scale", + "fc2_block_scale", + ) + _expected_num_inputs = 8 + + def _pack_weight_inputs( + self, view: dict[str, torch.Tensor], hidden_size: int + ) -> List[torch.Tensor]: + from math import ceil + + w1, w2, w1_scale, w2_scale = (view[key] for key in self._required_weight_keys) + num_experts = self.config.routing.num_experts + intermediate_size = self.config.experts.intermediate_size + expected_w1 = (num_experts, 2 * intermediate_size, hidden_size) + expected_w2 = (num_experts, hidden_size, intermediate_size) + expected_s1 = ( + num_experts, + ceil(2 * intermediate_size / 128), + ceil(hidden_size / 128), + ) + expected_s2 = ( + num_experts, + ceil(hidden_size / 128), + ceil(intermediate_size / 128), + ) + if w1.dtype is not torch.float8_e4m3fn or w2.dtype is not torch.float8_e4m3fn: + raise TypeError("Cutlass FP8-block prepared weights must be float8_e4m3fn.") + if w1_scale.dtype is not torch.float32 or w2_scale.dtype is not torch.float32: + raise TypeError("Cutlass FP8-block scales must be float32.") + if (tuple(w1.shape), tuple(w2.shape)) != (expected_w1, expected_w2): + raise ValueError( + f"Cutlass FP8-block weight shapes {tuple(w1.shape)}/{tuple(w2.shape)} " + f"!= expected {expected_w1}/{expected_w2}." + ) + if (tuple(w1_scale.shape), tuple(w2_scale.shape)) != (expected_s1, expected_s2): + raise ValueError( + "Cutlass FP8-block scale shapes " + f"{tuple(w1_scale.shape)}/{tuple(w2_scale.shape)} != expected " + f"{expected_s1}/{expected_s2}." + ) + self._validate_weight_storage((w1, w2, w1_scale, w2_scale)) + return [w1, w2, w1_scale, w2_scale] + + def _quant_scales(self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: + return [inputs[6], inputs[7]] + + +class CutlassMxfp8Mxfp4Runner(_CutlassRunnerBase): + """Unified adapter for CUTLASS MXFP8 x MXFP4 fused MoE.""" + + backend_key = "cutlass_mxfp8_mxfp4" + supported_quant_variants = (QuantVariant.MXFP4,) + _supported_archs = _CUTLASS_MXFP8_MXFP4_ARCHS + _x_dtype = torch.float8_e4m3fn + _weight_dtype = torch.int64 + _use_mxfp8_act_scaling = True + _required_weight_keys = ( + "fc1_expert_weights", + "fc2_expert_weights", + "fc1_expert_scales", + "fc2_expert_scales", + "fc1_input_scale", + "fc2_input_scale", + ) + _expected_num_inputs = 11 + + def _pack_weight_inputs( + self, view: dict[str, torch.Tensor], hidden_size: int + ) -> List[torch.Tensor]: + w1, w2, w1_scale, w2_scale, a1_scale, a2_scale = ( + view[key] for key in self._required_weight_keys + ) + num_experts = self.config.routing.num_experts + intermediate_size = self.config.experts.intermediate_size + if hidden_size % 128 != 0 or intermediate_size % 128 != 0: + raise ValueError( + "Cutlass MXFP8xMXFP4 requires hidden_size and intermediate_size " + f"divisible by 128, got H={hidden_size}, I={intermediate_size}." + ) + expected_w1 = (num_experts, 2 * intermediate_size, hidden_size // 2) + expected_w2 = (num_experts, hidden_size, intermediate_size // 2) + if w1.dtype is not torch.uint8 or w2.dtype is not torch.uint8: + raise TypeError("Cutlass MXFP8xMXFP4 packed weights must be uint8.") + if (tuple(w1.shape), tuple(w2.shape)) != (expected_w1, expected_w2): + raise ValueError( + "Cutlass MXFP8xMXFP4 weight shapes " + f"{tuple(w1.shape)}/{tuple(w2.shape)} != expected " + f"{expected_w1}/{expected_w2}." + ) + self._validate_weight_storage((w1, w2, w1_scale, w2_scale, a1_scale, a2_scale)) + expected_s1 = num_experts * _mxfp8_swizzled_act_sf_numel( + 2 * intermediate_size, hidden_size + ) + expected_s2 = num_experts * _mxfp8_swizzled_act_sf_numel( + hidden_size, intermediate_size + ) + if w1_scale.dtype is not torch.uint8 or w2_scale.dtype is not torch.uint8: + raise TypeError("Cutlass MXFP8xMXFP4 weight scales must be uint8.") + if w1_scale.numel() != expected_s1 or w2_scale.numel() != expected_s2: + raise ValueError( + "Cutlass MXFP8xMXFP4 weight scale sizes " + f"{w1_scale.numel()}/{w2_scale.numel()} != expected " + f"{expected_s1}/{expected_s2}." + ) + _require_cutlass_tensor( + a1_scale, + name="fc1_input_scale", + dtype=torch.float32, + shape=(num_experts,), + ) + _require_cutlass_tensor( + a2_scale, + name="fc2_input_scale", + dtype=torch.float32, + shape=(num_experts,), + ) + return [ + w1.view(torch.int64), + w2.view(torch.int64), + w1_scale, + w2_scale, + a1_scale, + a2_scale, + ] + + def _quant_scales(self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: + return [ + inputs[6].view(torch.int32), + inputs[8], + inputs[7].view(torch.int32), + inputs[9], + ] + + +class CutlassMxfp8Runner(_CutlassRunnerBase): + """Unified adapter for CUTLASS MXFP8 x MXFP8 fused MoE.""" + + backend_key = "cutlass_mxfp8" + supported_quant_variants = (QuantVariant.MxFp8,) + _supported_archs = _CUTLASS_MXFP8_ARCHS + _x_dtype = torch.float8_e4m3fn + _weight_dtype = torch.float8_e4m3fn + _use_mxfp8_act_scaling = True + _required_weight_keys = ( + "fc1_expert_weights", + "fc2_expert_weights", + "fc1_expert_scales", + "fc2_expert_scales", + "fc1_input_scale", + "fc2_input_scale", + ) + _expected_num_inputs = 11 + + def _pack_weight_inputs( + self, view: dict[str, torch.Tensor], hidden_size: int + ) -> List[torch.Tensor]: + w1, w2, w1_scale, w2_scale, a1_scale, a2_scale = ( + view[key] for key in self._required_weight_keys + ) + num_experts = self.config.routing.num_experts + intermediate_size = self.config.experts.intermediate_size + if hidden_size % 128 != 0 or intermediate_size % 128 != 0: + raise ValueError( + "Cutlass MXFP8 requires hidden_size and intermediate_size " + f"divisible by 128, got H={hidden_size}, I={intermediate_size}." + ) + expected_w1 = (num_experts, 2 * intermediate_size, hidden_size) + expected_w2 = (num_experts, hidden_size, intermediate_size) + if w1.dtype is not torch.float8_e4m3fn or w2.dtype is not torch.float8_e4m3fn: + raise TypeError("Cutlass MXFP8 prepared weights must be float8_e4m3fn.") + if (tuple(w1.shape), tuple(w2.shape)) != (expected_w1, expected_w2): + raise ValueError( + f"Cutlass MXFP8 weight shapes {tuple(w1.shape)}/{tuple(w2.shape)} " + f"!= expected {expected_w1}/{expected_w2}." + ) + # Binding uses alignToSfDim(I, 128) * 2 for gated SwiGLU, not + # round_up(2*I, 128). Those agree only when I % 128 == 0. + expected_s1 = ( + num_experts, + 2 * round_up(intermediate_size, 128), + round_up(hidden_size // 32, 4) // 4, + ) + expected_s2 = ( + num_experts, + round_up(hidden_size, 128), + round_up(intermediate_size // 32, 4) // 4, + ) + _require_cutlass_tensor( + w1_scale, name="fc1_expert_scales", dtype=torch.int32, shape=expected_s1 + ) + _require_cutlass_tensor( + w2_scale, name="fc2_expert_scales", dtype=torch.int32, shape=expected_s2 + ) + _require_cutlass_tensor( + a1_scale, + name="fc1_input_scale", + dtype=torch.float32, + shape=(num_experts,), + ) + _require_cutlass_tensor( + a2_scale, + name="fc2_input_scale", + dtype=torch.float32, + shape=(num_experts,), + ) + self._validate_weight_storage((w1, w2, w1_scale, w2_scale, a1_scale, a2_scale)) + return [w1, w2, w1_scale, w2_scale, a1_scale, a2_scale] + + def _quant_scales(self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: + return [inputs[6], inputs[8], inputs[7], inputs[9]] + + +class CutlassW4A8Runner(_CutlassRunnerBase): + """Unified adapter for CUTLASS INT4-weight x FP8-activation fused MoE.""" + + backend_key = "cutlass_w4a8" + supported_quant_variants = (QuantVariant.W4A8,) + _supported_archs = _CUTLASS_W4A8_ARCHS + _weight_dtype = torch.uint8 + _use_w4_group_scaling = True + _use_packed_weights = True + _required_weight_keys = ( + "fc1_expert_weights", + "fc2_expert_weights", + "fc1_expert_scales", + "fc2_expert_scales", + "fc1_act_scale", + "fc2_act_scale", + "fc1_zero", + "fc2_zero", + "fc1_alpha", + "fc2_alpha", + ) + _expected_num_inputs = 14 + + def _pack_weight_inputs( + self, view: dict[str, torch.Tensor], hidden_size: int + ) -> List[torch.Tensor]: + tensors = tuple(view[key] for key in self._required_weight_keys) + w1, w2 = tensors[0], tensors[1] + num_experts = self.config.routing.num_experts + intermediate_size = self.config.experts.intermediate_size + expected_w1 = (num_experts, 2 * intermediate_size, hidden_size // 2) + expected_w2 = (num_experts, hidden_size, intermediate_size // 2) + if w1.dtype is not torch.uint8 or w2.dtype is not torch.uint8: + raise TypeError("Cutlass W4A8 packed weights must be uint8.") + if (tuple(w1.shape), tuple(w2.shape)) != (expected_w1, expected_w2): + raise ValueError( + f"Cutlass W4A8 weight shapes {tuple(w1.shape)}/{tuple(w2.shape)} " + f"!= expected {expected_w1}/{expected_w2}." + ) + expected_s1 = ( + num_experts, + 2 * intermediate_size // 64, + hidden_size // 128, + 8, + 8, + ) + expected_s2 = ( + num_experts, + hidden_size // 64, + intermediate_size // 128, + 8, + 8, + ) + w1_scale, w2_scale = tensors[2], tensors[3] + act1, act2, zero1, zero2, alpha1, alpha2 = tensors[4:] + _require_cutlass_tensor( + w1_scale, + name="fc1_expert_scales", + dtype=torch.bfloat16, + shape=expected_s1, + ) + _require_cutlass_tensor( + w2_scale, + name="fc2_expert_scales", + dtype=torch.bfloat16, + shape=expected_s2, + ) + _require_cutlass_tensor( + act1, name="fc1_act_scale", dtype=torch.bfloat16, shape=(hidden_size,) + ) + _require_cutlass_tensor( + act2, + name="fc2_act_scale", + dtype=torch.bfloat16, + shape=(intermediate_size,), + ) + _require_cutlass_tensor( + zero1, name="fc1_zero", dtype=torch.bfloat16, shape=(0,) + ) + _require_cutlass_tensor( + zero2, name="fc2_zero", dtype=torch.bfloat16, shape=(0,) + ) + _require_cutlass_tensor( + alpha1, name="fc1_alpha", dtype=torch.float32, shape=(num_experts,) + ) + _require_cutlass_tensor( + alpha2, name="fc2_alpha", dtype=torch.float32, shape=(num_experts,) + ) + self._validate_weight_storage(tensors) + return list(tensors) + + def _quant_scales(self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: + return list(inputs[6:14]) + + +class CutlassHummingRunner(_CutlassRunnerBase): + """Unified adapter for CUTLASS Humming MXFP4 x FP8 fused MoE.""" + + backend_key = "cutlass_humming" + supported_quant_variants = (QuantVariant.Humming,) + _supported_archs = _CUTLASS_HUMMING_ARCHS + _weight_dtype = torch.uint8 + _use_w4_group_scaling = True + _use_wfp4afp8_humming = True + _required_weight_keys = ( + "fc1_expert_weights", + "fc2_expert_weights", + "fc1_expert_scales", + "fc2_expert_scales", + "fc1_residual_scale", + "fc2_residual_scale", + "fc2_act_global", + ) + _expected_num_inputs = 11 + + def _pack_weight_inputs( + self, view: dict[str, torch.Tensor], hidden_size: int + ) -> List[torch.Tensor]: + w1, w2, w1_scale, w2_scale, r1, r2, a2 = ( + view[key] for key in self._required_weight_keys + ) + num_experts = self.config.routing.num_experts + intermediate_size = self.config.experts.intermediate_size + expected_w1 = (num_experts, 2 * intermediate_size, hidden_size // 2) + expected_w2 = (num_experts, hidden_size, intermediate_size // 2) + if w1.dtype is not torch.uint8 or w2.dtype is not torch.uint8: + raise TypeError("Cutlass Humming packed weights must be uint8.") + if (tuple(w1.shape), tuple(w2.shape)) != (expected_w1, expected_w2): + raise ValueError( + "Cutlass Humming weight shapes " + f"{tuple(w1.shape)}/{tuple(w2.shape)} != expected " + f"{expected_w1}/{expected_w2}." + ) + expected_s1 = ( + num_experts, + 2 * intermediate_size // 64, + hidden_size // 128, + 16, + 16, + ) + expected_s2 = ( + num_experts, + hidden_size // 64, + intermediate_size // 128, + 16, + 16, + ) + _require_cutlass_tensor( + w1_scale, name="fc1_expert_scales", dtype=torch.uint8, shape=expected_s1 + ) + _require_cutlass_tensor( + w2_scale, name="fc2_expert_scales", dtype=torch.uint8, shape=expected_s2 + ) + _require_cutlass_tensor( + r1, name="fc1_residual_scale", dtype=torch.float32, shape=(num_experts,) + ) + _require_cutlass_tensor( + r2, name="fc2_residual_scale", dtype=torch.float32, shape=(num_experts,) + ) + _require_cutlass_tensor( + a2, name="fc2_act_global", dtype=torch.float32, shape=() + ) + self._validate_weight_storage((w1, w2, w1_scale, w2_scale, r1, r2, a2)) + return [w1, w2, w1_scale, w2_scale, r1, r2, a2] + + def _quant_scales(self, inputs: List[torch.Tensor]) -> List[torch.Tensor]: + return [ + inputs[6].view(torch.int32), + inputs[8], + inputs[10], + inputs[7].view(torch.int32), + inputs[9], + ] + + # --------------------------------------------------------------------------- # CuteDSL NVFP4 runner — delegates to the matching W4A4 or W4A16 runner # --------------------------------------------------------------------------- diff --git a/tests/moe/test_unified_moe.py b/tests/moe/test_unified_moe.py index 280d9a0f279..e6efa5ca28e 100644 --- a/tests/moe/test_unified_moe.py +++ b/tests/moe/test_unified_moe.py @@ -46,8 +46,11 @@ BackendOptions, CuteDslConfig, CuteDslNvfp4Runner, - CutlassConfig, CutlassBf16Config, + CutlassFp8BlockConfig, + CutlassFp8PerTensorConfig, + CutlassMxfp8Config, + CutlassNvfp4Config, ExecutionConfig, MoEFinalizeConfig, ExpertConfig, @@ -302,11 +305,11 @@ def test_finalize_config_custom(self): assert _eval_repr(cfg) == cfg def test_backend_options_multi(self): - opts = BackendOptions(candidates=(TrtllmFp4Config(), CutlassConfig())) + opts = BackendOptions(candidates=(TrtllmFp4Config(), CutlassNvfp4Config())) reconstructed = _eval_repr(opts) assert len(reconstructed) == 2 assert isinstance(reconstructed.candidates[0], TrtllmFp4Config) - assert isinstance(reconstructed.candidates[1], CutlassConfig) + assert isinstance(reconstructed.candidates[1], CutlassNvfp4Config) def test_backend_options_single(self): opts = BackendOptions(candidates=(TrtllmFp8PerTensorConfig(),)) @@ -336,7 +339,7 @@ def test_moe_config_full(self): experts=ExpertConfig(intermediate_size=2048, local_num_experts=32), activation=ActivationConfig(type=ActivationType.Geglu), backend=BackendOptions( - candidates=(TrtllmFp8BlockConfig(), CutlassConfig()) + candidates=(TrtllmFp8BlockConfig(), CutlassMxfp8Config()) ), execution=ExecutionConfig(enable_pdl=True, tune_max_num_tokens=4096), ) @@ -350,13 +353,13 @@ def test_moe_config_full(self): class TestBackendOptions: def test_explicit_candidates(self): - opts = BackendOptions(candidates=(TrtllmFp4Config(), CutlassConfig())) + opts = BackendOptions(candidates=(TrtllmFp4Config(), CutlassNvfp4Config())) assert isinstance(opts, BackendOptions) assert len(opts) == 2 def test_multiple_candidates(self): opts = BackendOptions( - candidates=(TrtllmFp4Config(), TrtllmFp8BlockConfig(), CutlassConfig()) + candidates=(TrtllmFp4Config(), TrtllmFp8BlockConfig(), CutlassNvfp4Config()) ) assert len(opts) == 3 @@ -375,7 +378,6 @@ def test_valid_for_blackwell(self): ) valid = opts.valid_for(100) assert len(valid) == 3 - assert not CutlassConfig.supported(100) assert CutlassBf16Config.supported(100) assert TrtllmBf16Config.supported(100) assert TrtllmBf16Config.supported(103) @@ -394,11 +396,11 @@ def test_valid_for_blackwell(self): assert not TrtllmFp8PerTensorConfig.supported(120) def test_iteration(self): - opts = BackendOptions(candidates=(TrtllmFp4Config(), CutlassConfig())) + opts = BackendOptions(candidates=(TrtllmFp4Config(), CutlassNvfp4Config())) items = list(opts) assert len(items) == 2 assert any(isinstance(c, TrtllmFp4Config) for c in items) - assert any(isinstance(c, CutlassConfig) for c in items) + assert any(isinstance(c, CutlassNvfp4Config) for c in items) def test_empty(self): opts = BackendOptions() @@ -488,7 +490,9 @@ def test_replace_backend(self): quant=QuantConfig(variant=QuantVariant.NVFP4), experts=ExpertConfig(intermediate_size=512), ) - narrow = dataclasses.replace(cfg, backend=BackendOptions((CutlassConfig(),))) + narrow = dataclasses.replace( + cfg, backend=BackendOptions((CutlassNvfp4Config(),)) + ) assert len(narrow.backend) == 1 @@ -570,7 +574,9 @@ def test_trtllm_fp4_deepseekv3(self): quant=QuantConfig(variant=QuantVariant.NVFP4), experts=ExpertConfig(intermediate_size=1024), activation=ActivationConfig(type=ActivationType.Swiglu), - backend=BackendOptions(candidates=(TrtllmFp4Config(), CutlassConfig())), + backend=BackendOptions( + candidates=(TrtllmFp4Config(), CutlassNvfp4Config()) + ), ) assert cfg.routing.method == RoutingMethodType.DeepSeekV3 assert cfg.quant.variant == QuantVariant.NVFP4 @@ -588,7 +594,7 @@ def test_trtllm_fp8_block_mxfp8(self): experts=ExpertConfig(intermediate_size=512), activation=ActivationConfig(type=ActivationType.Swiglu), backend=BackendOptions( - candidates=(TrtllmFp8BlockConfig(), CutlassConfig()) + candidates=(TrtllmFp8BlockConfig(), CutlassMxfp8Config()) ), ) assert cfg.quant.variant == QuantVariant.MxFp8 @@ -599,7 +605,9 @@ def test_trtllm_fp8_per_tensor(self): routing=RoutingConfig(num_experts=8, top_k=2), quant=QuantConfig(variant=QuantVariant.FP8PerTensor), experts=ExpertConfig(intermediate_size=512), - backend=BackendOptions((TrtllmFp8PerTensorConfig(),)), + backend=BackendOptions( + candidates=(TrtllmFp8PerTensorConfig(), CutlassFp8PerTensorConfig()) + ), ) assert cfg.quant.variant == QuantVariant.FP8PerTensor @@ -630,17 +638,15 @@ def test_trtllm_mxint4(self): assert cfg.quant.variant == QuantVariant.MxInt4 def test_cutlass_modular_fp8(self): - """Legacy declarative CUTLASS modular FP8 config.""" + """CUTLASS DeepSeek block-scale FP8 config.""" cfg = MoEConfig( routing=RoutingConfig(num_experts=64, top_k=8), quant=QuantConfig(variant=QuantVariant.DeepSeekFp8), experts=ExpertConfig(intermediate_size=2048), activation=ActivationConfig(type=ActivationType.Swiglu), - backend=BackendOptions((CutlassConfig(),)), + backend=BackendOptions((CutlassFp8BlockConfig(),)), ) - # CutlassConfig preserves the historical quant-neutral declarative form for - # compatibility; it is not a registered runnable backend. - assert any(isinstance(c, CutlassConfig) for c in cfg.backend) + assert any(isinstance(c, CutlassFp8BlockConfig) for c in cfg.backend) def test_cutedsl_nvfp4(self): """CuteDSL NVFP4 config.""" @@ -649,7 +655,7 @@ def test_cutedsl_nvfp4(self): quant=QuantConfig(variant=QuantVariant.NVFP4), experts=ExpertConfig(intermediate_size=1024), activation=ActivationConfig(type=ActivationType.Swiglu), - backend=BackendOptions(candidates=(CuteDslConfig(), CutlassConfig())), + backend=BackendOptions(candidates=(CuteDslConfig(), CutlassNvfp4Config())), ) assert any(isinstance(c, CuteDslConfig) for c in cfg.backend) diff --git a/tests/moe/test_unified_moe_cutlass.py b/tests/moe/test_unified_moe_cutlass.py index da561625f67..78f6b8ee8b5 100644 --- a/tests/moe/test_unified_moe_cutlass.py +++ b/tests/moe/test_unified_moe_cutlass.py @@ -1,4 +1,4 @@ -"""Unified CUTLASS BF16 and W4A16/MXFP4 MoE adapter tests.""" +"""Unified CUTLASS MoE adapter tests covering every quant-specific runner.""" from __future__ import annotations @@ -10,11 +10,24 @@ from flashinfer.fused_moe import ( ActivationConfig, BackendOptions, - CutlassBf16Runner, - CutlassConfig, CutlassBf16Config, + CutlassBf16Runner, + CutlassFp8BlockConfig, + CutlassFp8BlockRunner, + CutlassFp8PerTensorConfig, + CutlassFp8PerTensorRunner, + CutlassHummingConfig, + CutlassHummingRunner, + CutlassMxfp8Config, + CutlassMxfp8Mxfp4Config, + CutlassMxfp8Mxfp4Runner, + CutlassMxfp8Runner, + CutlassNvfp4Config, + CutlassNvfp4Runner, CutlassW4A16Config, CutlassW4A16Runner, + CutlassW4A8Config, + CutlassW4A8Runner, ExecutionConfig, ExpertConfig, MoEActivationPack, @@ -28,7 +41,7 @@ RoutingInputMode, ) from flashinfer.fused_moe.layer import _BACKEND_RUNNERS -from flashinfer.fused_moe.runners import MoERunner +from flashinfer.fused_moe.runners import MoERunner, _mxfp8_swizzled_act_sf_numel from flashinfer.fused_moe.prepare import _quantize_mxfp4_linear from flashinfer.fused_moe.utils import map_to_hybrid_bucket from flashinfer.tllm_enums import ActivationType @@ -39,6 +52,7 @@ is_sm110a_supported, is_sm120a_supported, is_sm121a_supported, + is_sm12x_supported, is_sm90a_supported, ) @@ -59,16 +73,34 @@ def _config(**overrides) -> MoEConfig: def test_cutlass_bf16_config_architectures_and_registration(): for arch in (89, 90, 100, 103, 107, 110, 120, 121): assert CutlassBf16Config.supported(arch) + assert CutlassFp8PerTensorConfig.supported(arch) assert CutlassW4A16Config.supported(90) + assert CutlassFp8BlockConfig.supported(90) + assert CutlassW4A8Config.supported(90) + assert CutlassHummingConfig.supported(90) + for arch in (100, 103, 107, 110, 120, 121): + assert CutlassNvfp4Config.supported(arch) + assert CutlassMxfp8Mxfp4Config.supported(arch) + for arch in (100, 103, 107): + assert CutlassMxfp8Config.supported(arch) assert not CutlassBf16Config.supported(80) assert not CutlassW4A16Config.supported(100) + assert not CutlassNvfp4Config.supported(90) + assert not CutlassFp8BlockConfig.supported(100) + assert not CutlassMxfp8Config.supported(90) + assert not CutlassMxfp8Config.supported(110) + assert not CutlassW4A8Config.supported(100) + assert not CutlassHummingConfig.supported(100) assert not CutlassBf16Config.supported(130) - assert CutlassConfig is not CutlassBf16Config - assert not CutlassConfig.supported(90) - assert BackendOptions((CutlassConfig(),)).valid_for(90) == [] - assert CutlassConfig not in _BACKEND_RUNNERS assert _BACKEND_RUNNERS[CutlassBf16Config] is CutlassBf16Runner + assert _BACKEND_RUNNERS[CutlassNvfp4Config] is CutlassNvfp4Runner assert _BACKEND_RUNNERS[CutlassW4A16Config] is CutlassW4A16Runner + assert _BACKEND_RUNNERS[CutlassFp8PerTensorConfig] is CutlassFp8PerTensorRunner + assert _BACKEND_RUNNERS[CutlassFp8BlockConfig] is CutlassFp8BlockRunner + assert _BACKEND_RUNNERS[CutlassMxfp8Mxfp4Config] is CutlassMxfp8Mxfp4Runner + assert _BACKEND_RUNNERS[CutlassMxfp8Config] is CutlassMxfp8Runner + assert _BACKEND_RUNNERS[CutlassW4A8Config] is CutlassW4A8Runner + assert _BACKEND_RUNNERS[CutlassHummingConfig] is CutlassHummingRunner def test_all_registered_runners_use_enforced_lifecycle(): @@ -151,11 +183,6 @@ def forward(self, inputs, **kwargs): runner.build() -def test_legacy_cutlass_config_is_deprecated(): - with pytest.warns(DeprecationWarning, match="CutlassConfig is deprecated"): - CutlassConfig() - - def test_prepare_cutlass_bf16_weights_preserves_canonical_layout(): w1 = torch.randn(2, 64, 64, dtype=torch.bfloat16)[..., ::2] w2 = torch.randn(2, 32, 64, dtype=torch.bfloat16)[..., ::2] @@ -199,6 +226,172 @@ def test_prepare_cutlass_w4a16_weights_rejects_invalid_source_contract(): ) +def test_prepare_cutlass_nvfp4_weights_rejects_invalid_source_contract(): + w1 = torch.empty(2, 512, 128, dtype=torch.float16) + w2 = torch.empty(2, 128, 256, dtype=torch.bfloat16) + with pytest.raises(TypeError, match="expects BF16 weights"): + CutlassNvfp4Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=128, + intermediate_size=256, + ) + + w1 = torch.empty(2, 510, 128, dtype=torch.bfloat16) + w2 = torch.empty(2, 128, 255, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="divisible by 16"): + CutlassNvfp4Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=128, + intermediate_size=255, + ) + + w1 = torch.empty(2, 32, 16, dtype=torch.bfloat16) + w2 = torch.empty(2, 16, 16, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="requires CUDA"): + CutlassNvfp4Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=16, + intermediate_size=16, + device=torch.device("cpu"), + ) + + +def test_prepare_cutlass_fp8_per_tensor_weights_rejects_invalid_source_contract(): + w1 = torch.empty(2, 64, 32, dtype=torch.float16) + w2 = torch.empty(2, 32, 32, dtype=torch.bfloat16) + with pytest.raises(TypeError, match="expects BF16 weights"): + CutlassFp8PerTensorConfig.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=32, + intermediate_size=32, + ) + + +def test_prepare_cutlass_fp8_block_weights_rejects_invalid_source_contract(): + w1 = torch.empty(2, 256, 128, dtype=torch.bfloat16) + w2 = torch.empty(2, 128, 128, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="divisible by 128"): + CutlassFp8BlockConfig.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=128, + intermediate_size=127, + ) + + +def test_prepare_cutlass_mxfp8_mxfp4_weights_rejects_invalid_source_contract(): + w1 = torch.empty(2, 64, 32, dtype=torch.float16) + w2 = torch.empty(2, 32, 32, dtype=torch.bfloat16) + with pytest.raises(TypeError, match="expects BF16 weights"): + CutlassMxfp8Mxfp4Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=32, + intermediate_size=32, + ) + + w1 = torch.empty(2, 128, 64, dtype=torch.bfloat16) + w2 = torch.empty(2, 64, 64, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="divisible by 128"): + CutlassMxfp8Mxfp4Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=64, + intermediate_size=64, + ) + + w1 = torch.empty(2, 256, 128, dtype=torch.bfloat16) + w2 = torch.empty(2, 128, 128, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="requires CUDA"): + CutlassMxfp8Mxfp4Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=128, + intermediate_size=128, + device=torch.device("cpu"), + ) + + +def test_prepare_cutlass_mxfp8_weights_rejects_invalid_source_contract(): + w1 = torch.empty(2, 64, 32, dtype=torch.float16) + w2 = torch.empty(2, 32, 32, dtype=torch.bfloat16) + with pytest.raises(TypeError, match="expects BF16 weights"): + CutlassMxfp8Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=32, + intermediate_size=32, + ) + + w1 = torch.empty(2, 128, 64, dtype=torch.bfloat16) + w2 = torch.empty(2, 64, 64, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="divisible by 128"): + CutlassMxfp8Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=64, + intermediate_size=64, + ) + + w1 = torch.empty(2, 384, 128, dtype=torch.bfloat16) + w2 = torch.empty(2, 128, 192, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="divisible by 128"): + CutlassMxfp8Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=128, + intermediate_size=192, + ) + + w1 = torch.empty(2, 256, 128, dtype=torch.bfloat16) + w2 = torch.empty(2, 128, 128, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="requires CUDA"): + CutlassMxfp8Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=128, + intermediate_size=128, + device=torch.device("cpu"), + ) + + +def test_prepare_cutlass_w4a8_and_humming_weights_reject_invalid_source_contract(): + w1 = torch.empty(2, 256, 128, dtype=torch.float16) + w2 = torch.empty(2, 128, 128, dtype=torch.bfloat16) + with pytest.raises(TypeError, match="expects BF16 weights"): + CutlassW4A8Config.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=128, + intermediate_size=128, + ) + with pytest.raises(TypeError, match="expects BF16 weights"): + CutlassHummingConfig.prepare_weights( + w1, + w2, + num_local_experts=2, + hidden_size=128, + intermediate_size=128, + ) + + def test_cutlass_mxfp4_linear_quantizer_code_points(): values = torch.tensor( [ @@ -336,6 +529,261 @@ def test_cutlass_runner_rejects_out_of_scope_configs(config, match): runner.check_support() +@pytest.mark.parametrize( + "config,match", + ( + ( + _config(quant=QuantConfig(variant=QuantVariant.BF16)), + "QuantVariant.BF16", + ), + ( + _config( + quant=QuantConfig(variant=QuantVariant.NVFP4), + activation=ActivationConfig(ActivationType.Relu2), + ), + "Swiglu", + ), + ( + _config( + quant=QuantConfig(variant=QuantVariant.NVFP4), + finalize=MoEFinalizeConfig(do_finalize=False), + ), + "do_finalize=True", + ), + ( + _config( + quant=QuantConfig(variant=QuantVariant.NVFP4), + experts=ExpertConfig( + intermediate_size=256, + local_expert_offset=2, + local_num_experts=2, + ), + ), + "expert parallelism", + ), + ), +) +def test_cutlass_nvfp4_runner_rejects_out_of_scope_configs(config, match): + runner = CutlassNvfp4Runner.__new__(CutlassNvfp4Runner) + runner.config = config + with pytest.raises(NotImplementedError, match=match): + runner.check_support() + + +@pytest.mark.parametrize( + "runner_cls,quant,match", + ( + (CutlassFp8PerTensorRunner, QuantVariant.BF16, "QuantVariant.BF16"), + (CutlassFp8BlockRunner, QuantVariant.BF16, "QuantVariant.BF16"), + (CutlassMxfp8Mxfp4Runner, QuantVariant.BF16, "QuantVariant.BF16"), + (CutlassMxfp8Runner, QuantVariant.BF16, "QuantVariant.BF16"), + (CutlassW4A8Runner, QuantVariant.BF16, "QuantVariant.BF16"), + (CutlassHummingRunner, QuantVariant.BF16, "QuantVariant.BF16"), + ( + CutlassFp8PerTensorRunner, + QuantVariant.FP8PerTensor, + "do_finalize=True", + ), + ), +) +def test_cutlass_quant_runners_reject_out_of_scope_configs(runner_cls, quant, match): + if match == "do_finalize=True": + config = _config( + quant=QuantConfig(variant=quant), + finalize=MoEFinalizeConfig(do_finalize=False), + ) + else: + config = _config(quant=QuantConfig(variant=quant)) + runner = runner_cls.__new__(runner_cls) + runner.config = config + with pytest.raises(NotImplementedError, match=match): + runner.check_support() + + +def test_cutlass_mxfp8_rejects_linear_scale_layout(): + runner = CutlassMxfp8Runner.__new__(CutlassMxfp8Runner) + runner.config = _config( + quant=QuantConfig(variant=QuantVariant.MxFp8, swizzled_scale_factors=False) + ) + runner._device_arch = 100 + with pytest.raises(NotImplementedError, match="swizzled MXFP8 input_sf"): + runner.check_support() + + +def test_cutlass_mxfp8_mxfp4_rejects_linear_scale_layout(): + runner = CutlassMxfp8Mxfp4Runner.__new__(CutlassMxfp8Mxfp4Runner) + runner.config = _config( + quant=QuantConfig(variant=QuantVariant.MXFP4, swizzled_scale_factors=False) + ) + runner._device_arch = 100 + with pytest.raises(NotImplementedError, match="swizzled MXFP8 input_sf"): + runner.check_support() + + +def test_cutlass_fp8_block_rejects_cuda_below_12_8(monkeypatch): + monkeypatch.setattr( + "flashinfer.jit.cpp_ext.is_cuda_version_at_least", + lambda _version: False, + ) + runner = CutlassFp8BlockRunner.__new__(CutlassFp8BlockRunner) + runner.config = _config(quant=QuantConfig(variant=QuantVariant.DeepSeekFp8)) + runner._device_arch = 90 + with pytest.raises(NotImplementedError, match="requires CUDA 12.8 or newer"): + runner.check_support() + + +def test_cutlass_mxfp8_rejects_linear_activation_scales(): + runner = CutlassMxfp8Runner.__new__(CutlassMxfp8Runner) + hidden = torch.empty(16, 128, dtype=torch.float8_e4m3fn) + linear_sf = torch.empty(16, 4, dtype=torch.uint8) + act = MoEActivationPack( + hidden, + linear_sf, + torch.zeros(16, 2, dtype=torch.int32), + torch.ones(16, 2, dtype=torch.float32), + ) + with pytest.raises(ValueError, match="swizzled"): + runner._validate_activation_scale(act) + + +def test_cutlass_fp8_per_tensor_rejects_nonscalar_activation_scale(): + runner = CutlassFp8PerTensorRunner.__new__(CutlassFp8PerTensorRunner) + hidden = torch.empty(1, 128, dtype=torch.float8_e4m3fn) + act = MoEActivationPack( + hidden, + torch.ones(1, dtype=torch.float32), + torch.zeros(1, 2, dtype=torch.int32), + torch.ones(1, 2, dtype=torch.float32), + ) + with pytest.raises(ValueError, match="0-dim float32"): + runner._validate_activation_scale(act) + + +def test_cutlass_fp8_per_tensor_pack_keeps_scalar_scale_static(): + runner = CutlassFp8PerTensorRunner.__new__(CutlassFp8PerTensorRunner) + runner.config = _config(quant=QuantConfig(variant=QuantVariant.FP8PerTensor)) + runner.device = torch.device("cpu") + runner._inner = object() + runner._built = True + runner._ensure_workspace = lambda *_args, **_kwargs: None + runner._pack_weight_inputs = lambda _view, _hidden_size: [ + torch.empty(1) for _ in runner._required_weight_keys + ] + act = MoEActivationPack( + torch.empty(1, 128, dtype=torch.float8_e4m3fn), + torch.ones((), dtype=torch.float32), + torch.zeros(1, 2, dtype=torch.int32), + torch.full((1, 2), 0.5, dtype=torch.float32), + ) + weights = MoEWeightPack() + weights.prepare_for( + runner.backend_key, + {key: torch.empty(1) for key in runner._required_weight_keys}, + ) + runner.pack_inputs(act, weights) + spec = runner.tuning_config.dynamic_tensor_specs[0] + assert spec.input_idx == (0, 1, 2, 3) + + +def test_cutlass_mxfp8_mxfp4_pack_rejects_unaligned_hidden_size(): + runner = CutlassMxfp8Mxfp4Runner.__new__(CutlassMxfp8Mxfp4Runner) + runner.config = _config( + quant=QuantConfig(variant=QuantVariant.MXFP4), + routing=RoutingConfig(num_experts=2, top_k=2), + experts=ExpertConfig(intermediate_size=64), + ) + runner.device = torch.device("cpu") + view = {key: torch.empty(1) for key in runner._required_weight_keys} + with pytest.raises(ValueError, match="divisible by 128"): + runner._pack_weight_inputs(view, hidden_size=64) + + +def test_cutlass_mxfp8_pack_rejects_unaligned_hidden_size(): + runner = CutlassMxfp8Runner.__new__(CutlassMxfp8Runner) + runner.config = _config( + quant=QuantConfig(variant=QuantVariant.MxFp8), + routing=RoutingConfig(num_experts=2, top_k=2), + experts=ExpertConfig(intermediate_size=64), + ) + runner.device = torch.device("cpu") + view = {key: torch.empty(1) for key in runner._required_weight_keys} + with pytest.raises(ValueError, match="divisible by 128"): + runner._pack_weight_inputs(view, hidden_size=64) + + runner.config = _config( + quant=QuantConfig(variant=QuantVariant.MxFp8), + routing=RoutingConfig(num_experts=2, top_k=2), + experts=ExpertConfig(intermediate_size=192), + ) + with pytest.raises(ValueError, match="divisible by 128"): + runner._pack_weight_inputs(view, hidden_size=128) + + +def test_cutlass_mxfp8_pack_rejects_malformed_weight_scales(): + runner = CutlassMxfp8Runner.__new__(CutlassMxfp8Runner) + runner.config = _config( + quant=QuantConfig(variant=QuantVariant.MxFp8), + routing=RoutingConfig(num_experts=2, top_k=2), + experts=ExpertConfig(intermediate_size=256), + ) + runner.device = torch.device("cpu") + view = { + "fc1_expert_weights": torch.empty(2, 512, 128, dtype=torch.float8_e4m3fn), + "fc2_expert_weights": torch.empty(2, 128, 256, dtype=torch.float8_e4m3fn), + "fc1_expert_scales": torch.empty(2, 4, dtype=torch.int32), + "fc2_expert_scales": torch.empty(2, 4, dtype=torch.int32), + "fc1_input_scale": torch.ones(2, dtype=torch.float32), + "fc2_input_scale": torch.ones(2, dtype=torch.float32), + } + with pytest.raises(ValueError, match="fc1_expert_scales"): + runner._pack_weight_inputs(view, hidden_size=128) + + +def test_cutlass_w4a8_pack_rejects_malformed_weight_scales(): + runner = CutlassW4A8Runner.__new__(CutlassW4A8Runner) + runner.config = _config( + quant=QuantConfig(variant=QuantVariant.W4A8), + routing=RoutingConfig(num_experts=2, top_k=2), + experts=ExpertConfig(intermediate_size=256), + ) + runner.device = torch.device("cpu") + view = { + "fc1_expert_weights": torch.empty(2, 512, 64, dtype=torch.uint8), + "fc2_expert_weights": torch.empty(2, 128, 128, dtype=torch.uint8), + "fc1_expert_scales": torch.empty(2, 4, dtype=torch.bfloat16), + "fc2_expert_scales": torch.empty(2, 4, dtype=torch.bfloat16), + "fc1_act_scale": torch.ones(128, dtype=torch.bfloat16), + "fc2_act_scale": torch.ones(256, dtype=torch.bfloat16), + "fc1_zero": torch.empty(0, dtype=torch.bfloat16), + "fc2_zero": torch.empty(0, dtype=torch.bfloat16), + "fc1_alpha": torch.ones(2, dtype=torch.float32), + "fc2_alpha": torch.ones(2, dtype=torch.float32), + } + with pytest.raises(ValueError, match="fc1_expert_scales"): + runner._pack_weight_inputs(view, hidden_size=128) + + +def test_cutlass_humming_pack_rejects_malformed_weight_scales(): + runner = CutlassHummingRunner.__new__(CutlassHummingRunner) + runner.config = _config( + quant=QuantConfig(variant=QuantVariant.Humming), + routing=RoutingConfig(num_experts=2, top_k=2), + experts=ExpertConfig(intermediate_size=256), + ) + runner.device = torch.device("cpu") + view = { + "fc1_expert_weights": torch.empty(2, 512, 64, dtype=torch.uint8), + "fc2_expert_weights": torch.empty(2, 128, 128, dtype=torch.uint8), + "fc1_expert_scales": torch.empty(2, 4, dtype=torch.uint8), + "fc2_expert_scales": torch.empty(2, 4, dtype=torch.uint8), + "fc1_residual_scale": torch.ones(2, dtype=torch.float32), + "fc2_residual_scale": torch.ones(2, dtype=torch.float32), + "fc2_act_global": torch.ones((), dtype=torch.float32), + } + with pytest.raises(ValueError, match="fc1_expert_scales"): + runner._pack_weight_inputs(view, hidden_size=128) + + def test_moe_layer_checks_support_before_build_and_execution(monkeypatch): from flashinfer.fused_moe import layer as layer_module @@ -479,16 +927,6 @@ def test_cutlass_direct_execution_requires_explicit_build(monkeypatch, execute): assert backend_calls == [] -def test_legacy_cutlass_config_is_not_runnable(monkeypatch): - from flashinfer.fused_moe import layer as layer_module - - monkeypatch.setattr(layer_module, "get_compute_capability", lambda device: (9, 0)) - config = _config(backend=BackendOptions((CutlassConfig(),))) - - with pytest.raises(RuntimeError, match="none of the configured backends"): - MoELayer(config, device=torch.device("cuda")) - - def test_cutlass_autotuner_preparation_initializes_both_gemms(): class RecordingInner: def __init__(self): @@ -864,6 +1302,24 @@ def _reference(act: MoEActivationPack, w1: torch.Tensor, w2: torch.Tensor): return result.to(torch.bfloat16) +def _dequant_linear_mxfp4(packed: torch.Tensor, scales: torch.Tensor) -> torch.Tensor: + """Dequant packed E2M1 + linear UE8M0 scales without Humming preprocessing.""" + low = packed & 0xF + high = packed >> 4 + codes = torch.stack((low, high), dim=-1).reshape( + packed.shape[0], packed.shape[1], packed.shape[2] * 2 + ) + magnitudes = torch.tensor( + [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], + device=packed.device, + dtype=torch.float32, + ) + values = magnitudes[codes.to(torch.long) & 0x7] + values = torch.where((codes & 0x8) != 0, -values, values) + scale = torch.exp2(scales.to(torch.int16).to(torch.float32) - 127) + return values * scale.repeat_interleave(32, dim=-1) + + def _assert_numerically_close( actual: torch.Tensor, expected: torch.Tensor, @@ -1071,3 +1527,708 @@ def test_cutlass_w4a16_autotuned_compound_tactic_and_cuda_graph(): torch.cuda.synchronize() _assert_numerically_close(captured, expected, rtol=5e-2, atol=2e-2) + + +def _is_cutlass_nvfp4_runtime_supported() -> bool: + if not torch.cuda.is_available(): + return False + device = torch.device("cuda") + major, minor = get_compute_capability(device) + arch = major * 10 + minor + if not CutlassNvfp4Config.supported(arch): + return False + if arch in (100, 103): + return is_sm100a_supported(device) + if arch == 107: + return is_sm100f_supported(device) + if arch == 110: + return is_sm110a_supported(device) + if arch in (120, 121): + return is_sm12x_supported(device) + return False + + +cutlass_nvfp4_required = pytest.mark.skipif( + not _is_cutlass_nvfp4_runtime_supported(), + reason="requires SM100/SM110/SM12x CUTLASS NVFP4 GPU and CUDA toolkit", +) + + +def _make_nvfp4_case(num_tokens: int = 16): + torch.manual_seed(44) + device = torch.device("cuda", torch.cuda.current_device()) + num_experts, top_k = 4, 2 + hidden_size, intermediate_size = 128, 256 + x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) / 2 + w1 = ( + torch.randn( + num_experts, + 2 * intermediate_size, + hidden_size, + device=device, + dtype=torch.bfloat16, + ) + / 10 + ) + w2 = ( + torch.randn( + num_experts, + hidden_size, + intermediate_size, + device=device, + dtype=torch.bfloat16, + ) + / 10 + ) + topk_ids = torch.stack( + [torch.randperm(num_experts, device=device)[:top_k] for _ in range(num_tokens)] + ).to(torch.int32) + topk_weights = torch.softmax(torch.randn(num_tokens, top_k, device=device), dim=-1) + config = _config( + quant=QuantConfig(variant=QuantVariant.NVFP4), + experts=ExpertConfig(intermediate_size=intermediate_size), + backend=BackendOptions((CutlassNvfp4Config(),)), + execution=ExecutionConfig( + enable_pdl=False, + tune_max_num_tokens=max(64, num_tokens), + ), + ) + act = MoEActivationPack(x, None, topk_ids, topk_weights) + view = CutlassNvfp4Config.prepare_weights( + w1, + w2, + num_local_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + weights = MoEWeightPack() + weights.prepare_for("cutlass_nvfp4", view) + return config, act, weights, view + + +def _dequantize_cutlass_nvfp4_matrix( + packed: torch.Tensor, scale: torch.Tensor +) -> torch.Tensor: + from flashinfer.fp4_quantization import e2m1_and_ufp8sf_scale_to_float + + rows, packed_cols = packed.shape + global_scale = torch.ones(1, dtype=torch.float32) + return e2m1_and_ufp8sf_scale_to_float( + packed, + scale.reshape(-1), + global_scale, + sf_vec_size=16, + ufp8_type=1, + is_sf_swizzled_layout=True, + ).view(rows, packed_cols * 2) + + +def _nvfp4_quantized_reference(act: MoEActivationPack, view: dict[str, torch.Tensor]): + from flashinfer.fp4_quantization import fp4_quantize + + x = act.hidden_states_q + global_scale = torch.ones(1, device=x.device, dtype=torch.float32) + x_q, x_sf = fp4_quantize( + x, + global_scale=global_scale, + sf_vec_size=16, + is_sf_swizzled_layout=True, + ) + x_dq = _dequantize_cutlass_nvfp4_matrix(x_q, x_sf).to( + device=x.device, dtype=torch.bfloat16 + ) + w1_q = view["fc1_expert_weights"] + w2_q = view["fc2_expert_weights"] + w1_sf = view["fc1_weight_block_scale"] + w2_sf = view["fc2_weight_block_scale"] + w1 = torch.stack( + [ + _dequantize_cutlass_nvfp4_matrix(w1_q[i], w1_sf[i]).to( + device=x.device, dtype=torch.bfloat16 + ) + for i in range(w1_q.shape[0]) + ] + ) + w2 = torch.stack( + [ + _dequantize_cutlass_nvfp4_matrix(w2_q[i], w2_sf[i]).to( + device=x.device, dtype=torch.bfloat16 + ) + for i in range(w2_q.shape[0]) + ] + ) + ref_act = MoEActivationPack(x_dq, None, act.topk_ids, act.topk_weights) + return _reference(ref_act, w1, w2) + + +@cutlass_nvfp4_required +def test_cutlass_nvfp4_moe_layer_matches_quantized_reference(): + config, act, weights, view = _make_nvfp4_case() + assert view["fc1_expert_weights"].dtype is torch.uint8 + assert view["fc1_weight_block_scale"].ndim == 3 + layer = MoELayer(config) + runner = _pin_fallback_winner(layer, act) + + actual = layer(act, weights) + expected = _nvfp4_quantized_reference(act, view) + + assert layer.winner_backend == "cutlass_nvfp4" + assert runner._workspace is not None + assert torch.isfinite(actual).all(), "CUTLASS NVFP4 produced non-finite output" + _assert_numerically_close(actual, expected, rtol=2e-1, atol=2e-1) + + +@cutlass_nvfp4_required +def test_cutlass_nvfp4_autotuned_compound_tactic_and_cuda_graph(): + config, act, weights, view = _make_nvfp4_case(num_tokens=17) + runner = MoELayer(config).runners[0] + inputs = runner.pack_inputs(act, weights) + assert inputs[4].dtype is torch.int64 + assert len(inputs) == 12 + + with autotune(True): + _, tactic = AutoTuner.get().choose_one( + "test_moe_cutlass_nvfp4_compound", + [runner], + runner.tuning_config, + inputs, + ) + assert isinstance(tactic, tuple) and len(tactic) == 2 + assert all(stage_tactic >= 0 for stage_tactic in tactic) + + actual = runner.forward(inputs, tactic=tactic) + torch.cuda.synchronize() + expected = _nvfp4_quantized_reference(act, view) + assert torch.isfinite(actual).all(), "CUTLASS NVFP4 produced non-finite output" + _assert_numerically_close(actual, expected, rtol=2e-1, atol=2e-1) + + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = runner.forward(inputs, tactic=tactic) + captured_workspace = runner._workspace + runner._ensure_workspace(64, inputs[1].shape[1]) + assert runner._workspace is not captured_workspace + assert runner._workspace_cache[(32, inputs[1].shape[1])] is captured_workspace + captured.fill_(float("nan")) + graph.replay() + torch.cuda.synchronize() + + _assert_numerically_close(captured, expected, rtol=2e-1, atol=2e-1) + + +def _is_cutlass_fp8_runtime_supported() -> bool: + return _is_cutlass_bf16_runtime_supported() + + +cutlass_fp8_required = pytest.mark.skipif( + not _is_cutlass_fp8_runtime_supported(), + reason="requires a supported CUTLASS FP8 GPU and CUDA toolkit", +) + +cutlass_fp8_block_required = pytest.mark.skipif( + not torch.cuda.is_available() or not is_sm90a_supported(torch.device("cuda")), + reason="requires SM90a CUTLASS DeepSeek FP8 block scaling", +) + + +def _is_cutlass_mxfp8_runtime_supported() -> bool: + if not torch.cuda.is_available(): + return False + device = torch.device("cuda") + major, minor = get_compute_capability(device) + arch = major * 10 + minor + if not CutlassMxfp8Config.supported(arch): + return False + if arch in (100, 103): + return is_sm100a_supported(device) + if arch == 107: + return is_sm100f_supported(device) + return False + + +cutlass_mxfp8_mxfp4_required = pytest.mark.skipif( + not _is_cutlass_nvfp4_runtime_supported(), + reason="requires SM100/SM110/SM12x CUTLASS MXFP8xMXFP4 GPU and CUDA toolkit", +) + +cutlass_mxfp8_required = pytest.mark.skipif( + not _is_cutlass_mxfp8_runtime_supported(), + reason="requires SM100/SM103/SM107 CUTLASS MXFP8xMXFP8 GPU and CUDA toolkit", +) + +cutlass_w4a8_required = pytest.mark.skipif( + not torch.cuda.is_available() or not is_sm90a_supported(torch.device("cuda")), + reason="requires SM90a CUTLASS W4A8", +) + +cutlass_humming_required = pytest.mark.skipif( + not torch.cuda.is_available() or not is_sm90a_supported(torch.device("cuda")), + reason="requires SM90a CUTLASS Humming", +) + + +def _make_routing(num_tokens, num_experts, top_k, device): + topk_ids = torch.stack( + [torch.randperm(num_experts, device=device)[:top_k] for _ in range(num_tokens)] + ).to(torch.int32) + topk_weights = torch.softmax(torch.randn(num_tokens, top_k, device=device), dim=-1) + return topk_ids, topk_weights + + +def _make_bf16_experts(num_experts, hidden_size, intermediate_size, device): + w1 = ( + torch.randn( + num_experts, + 2 * intermediate_size, + hidden_size, + device=device, + dtype=torch.bfloat16, + ) + / 10 + ) + w2 = ( + torch.randn( + num_experts, + hidden_size, + intermediate_size, + device=device, + dtype=torch.bfloat16, + ) + / 10 + ) + return w1, w2 + + +def _autotune_and_graph(runner, act, weights, expected, *, rtol, atol, cache_name): + inputs = runner.pack_inputs(act, weights) + with autotune(True): + _, tactic = AutoTuner.get().choose_one( + cache_name, + [runner], + runner.tuning_config, + inputs, + ) + assert isinstance(tactic, tuple) and len(tactic) == 2 + assert all(stage_tactic >= 0 for stage_tactic in tactic) + actual = runner.forward(inputs, tactic=tactic) + torch.cuda.synchronize() + assert torch.isfinite(actual).all() + _assert_numerically_close(actual, expected, rtol=rtol, atol=atol) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = runner.forward(inputs, tactic=tactic) + captured.fill_(float("nan")) + graph.replay() + torch.cuda.synchronize() + _assert_numerically_close(captured, expected, rtol=rtol, atol=atol) + + +@cutlass_fp8_required +def test_cutlass_fp8_per_tensor_moe_layer_matches_quantized_reference(): + torch.manual_seed(45) + device = torch.device("cuda", torch.cuda.current_device()) + num_tokens, num_experts, top_k = 16, 4, 2 + hidden_size, intermediate_size = 128, 256 + x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) / 2 + w1, w2 = _make_bf16_experts(num_experts, hidden_size, intermediate_size, device) + topk_ids, topk_weights = _make_routing(num_tokens, num_experts, top_k, device) + x_q, x_scale = CutlassFp8PerTensorConfig.prepare_activations(x) + view = CutlassFp8PerTensorConfig.prepare_weights( + w1, + w2, + num_local_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + x_dq = x_q.float() * x_scale + w1_dq = view["fc1_expert_weights"].float() * view["fc1_dequant"][:, None, None] + w2_dq = view["fc2_expert_weights"].float() * view["fc2_dequant"][:, None, None] + config = _config( + quant=QuantConfig(variant=QuantVariant.FP8PerTensor), + experts=ExpertConfig(intermediate_size=intermediate_size), + backend=BackendOptions((CutlassFp8PerTensorConfig(),)), + execution=ExecutionConfig(enable_pdl=False, tune_max_num_tokens=64), + ) + act = MoEActivationPack(x_q, x_scale, topk_ids, topk_weights) + weights = MoEWeightPack() + weights.prepare_for("cutlass_fp8_per_tensor", view) + layer = MoELayer(config) + _pin_fallback_winner(layer, act) + actual = layer(act, weights) + expected = _reference( + MoEActivationPack(x_dq.to(torch.bfloat16), None, topk_ids, topk_weights), + w1_dq.to(torch.bfloat16), + w2_dq.to(torch.bfloat16), + ) + assert layer.winner_backend == "cutlass_fp8_per_tensor" + _assert_numerically_close(actual, expected, rtol=1e-1, atol=1e-1) + _autotune_and_graph( + layer.runners[0], + act, + weights, + expected, + rtol=1e-1, + atol=1e-1, + cache_name="test_moe_cutlass_fp8_per_tensor", + ) + + +@cutlass_fp8_block_required +def test_cutlass_fp8_block_moe_layer_matches_quantized_reference(): + torch.manual_seed(46) + device = torch.device("cuda", torch.cuda.current_device()) + num_tokens, num_experts, top_k = 16, 4, 2 + hidden_size, intermediate_size = 128, 256 + x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) / 2 + w1, w2 = _make_bf16_experts(num_experts, hidden_size, intermediate_size, device) + topk_ids, topk_weights = _make_routing(num_tokens, num_experts, top_k, device) + view = CutlassFp8BlockConfig.prepare_weights( + w1, + w2, + num_local_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + w1_dq = view["fc1_expert_weights"].float() * view[ + "fc1_block_scale" + ].repeat_interleave(128, dim=-2).repeat_interleave(128, dim=-1) + w2_dq = view["fc2_expert_weights"].float() * view[ + "fc2_block_scale" + ].repeat_interleave(128, dim=-2).repeat_interleave(128, dim=-1) + config = _config( + quant=QuantConfig(variant=QuantVariant.DeepSeekFp8), + experts=ExpertConfig(intermediate_size=intermediate_size), + backend=BackendOptions((CutlassFp8BlockConfig(),)), + execution=ExecutionConfig(enable_pdl=False, tune_max_num_tokens=64), + ) + act = MoEActivationPack(x, None, topk_ids, topk_weights) + weights = MoEWeightPack() + weights.prepare_for("cutlass_fp8_block", view) + layer = MoELayer(config) + _pin_fallback_winner(layer, act) + actual = layer(act, weights) + expected = _reference(act, w1_dq.to(torch.bfloat16), w2_dq.to(torch.bfloat16)) + assert layer.winner_backend == "cutlass_fp8_block" + _assert_numerically_close(actual, expected, rtol=1e-1, atol=1e-1) + _autotune_and_graph( + layer.runners[0], + act, + weights, + expected, + rtol=1e-1, + atol=1e-1, + cache_name="test_moe_cutlass_fp8_block", + ) + + +@cutlass_mxfp8_mxfp4_required +def test_cutlass_mxfp8_mxfp4_moe_layer_matches_quantized_reference(): + from flashinfer import mxfp4_dequantize, mxfp8_dequantize_host + + torch.manual_seed(47) + device = torch.device("cuda", torch.cuda.current_device()) + num_tokens, num_experts, top_k = 16, 4, 2 + hidden_size, intermediate_size = 128, 256 + x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) / 2 + w1, w2 = _make_bf16_experts(num_experts, hidden_size, intermediate_size, device) + topk_ids, topk_weights = _make_routing(num_tokens, num_experts, top_k, device) + x_q, x_sf = CutlassMxfp8Mxfp4Config.prepare_activations(x) + view = CutlassMxfp8Mxfp4Config.prepare_weights( + w1, + w2, + num_local_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + x_dq = mxfp8_dequantize_host( + x_q.cpu().view(torch.uint8), + x_sf.cpu().view(torch.uint8).reshape(-1), + True, + ).to(device=device, dtype=torch.bfloat16) + w1_dq = torch.stack( + [ + mxfp4_dequantize( + view["fc1_expert_weights"][i].cpu(), + view["fc1_expert_scales"][i].cpu().view(torch.uint8).reshape(-1), + ) + for i in range(num_experts) + ] + ).to(device=device, dtype=torch.bfloat16) + w2_dq = torch.stack( + [ + mxfp4_dequantize( + view["fc2_expert_weights"][i].cpu(), + view["fc2_expert_scales"][i].cpu().view(torch.uint8).reshape(-1), + ) + for i in range(num_experts) + ] + ).to(device=device, dtype=torch.bfloat16) + config = _config( + quant=QuantConfig(variant=QuantVariant.MXFP4), + experts=ExpertConfig(intermediate_size=intermediate_size), + backend=BackendOptions((CutlassMxfp8Mxfp4Config(),)), + execution=ExecutionConfig(enable_pdl=False, tune_max_num_tokens=16), + ) + act = MoEActivationPack(x_q, x_sf, topk_ids, topk_weights) + weights = MoEWeightPack() + weights.prepare_for("cutlass_mxfp8_mxfp4", view) + layer = MoELayer(config) + _pin_fallback_winner(layer, act) + actual = layer(act, weights) + expected = _reference( + MoEActivationPack(x_dq, None, topk_ids, topk_weights), w1_dq, w2_dq + ) + assert layer.winner_backend == "cutlass_mxfp8_mxfp4" + _assert_numerically_close(actual, expected, rtol=1e-1, atol=1e-1) + _autotune_and_graph( + layer.runners[0], + act, + weights, + expected, + rtol=1e-1, + atol=1e-1, + cache_name="test_moe_cutlass_mxfp8_mxfp4", + ) + + +@cutlass_mxfp8_required +def test_cutlass_mxfp8_moe_layer_matches_quantized_reference(): + from flashinfer import mxfp8_dequantize_host + + torch.manual_seed(48) + device = torch.device("cuda", torch.cuda.current_device()) + num_tokens, num_experts, top_k = 16, 4, 2 + hidden_size, intermediate_size = 128, 256 + x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) / 2 + w1, w2 = _make_bf16_experts(num_experts, hidden_size, intermediate_size, device) + topk_ids, topk_weights = _make_routing(num_tokens, num_experts, top_k, device) + x_q, x_sf = CutlassMxfp8Config.prepare_activations(x) + view = CutlassMxfp8Config.prepare_weights( + w1, + w2, + num_local_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + x_dq = mxfp8_dequantize_host( + x_q.cpu().view(torch.uint8), + x_sf.cpu().view(torch.uint8).reshape(-1), + True, + ).to(device=device, dtype=torch.bfloat16) + config = _config( + quant=QuantConfig(variant=QuantVariant.MxFp8), + experts=ExpertConfig(intermediate_size=intermediate_size), + backend=BackendOptions((CutlassMxfp8Config(),)), + execution=ExecutionConfig(enable_pdl=False, tune_max_num_tokens=16), + ) + act = MoEActivationPack(x_q, x_sf, topk_ids, topk_weights) + weights = MoEWeightPack() + weights.prepare_for("cutlass_mxfp8", view) + layer = MoELayer(config) + _pin_fallback_winner(layer, act) + actual = layer(act, weights) + expected = _reference(MoEActivationPack(x_dq, None, topk_ids, topk_weights), w1, w2) + assert layer.winner_backend == "cutlass_mxfp8" + assert torch.isfinite(actual).all() + _assert_numerically_close(actual, expected, rtol=2e-1, atol=2e-1) + _autotune_and_graph( + layer.runners[0], + act, + weights, + expected, + rtol=2e-1, + atol=2e-1, + cache_name="test_moe_cutlass_mxfp8", + ) + + +@cutlass_mxfp8_required +def test_cutlass_mxfp8_autotune_regenerates_swizzled_input_sf_across_bucket(): + torch.manual_seed(51) + device = torch.device("cuda", torch.cuda.current_device()) + num_tokens, num_experts, top_k = 257, 4, 2 + hidden_size, intermediate_size = 128, 256 + x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) / 2 + w1, w2 = _make_bf16_experts(num_experts, hidden_size, intermediate_size, device) + topk_ids, topk_weights = _make_routing(num_tokens, num_experts, top_k, device) + x_q, x_sf = CutlassMxfp8Config.prepare_activations(x) + view = CutlassMxfp8Config.prepare_weights( + w1, + w2, + num_local_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + config = _config( + quant=QuantConfig(variant=QuantVariant.MxFp8), + experts=ExpertConfig(intermediate_size=intermediate_size), + backend=BackendOptions((CutlassMxfp8Config(),)), + execution=ExecutionConfig(enable_pdl=False, tune_max_num_tokens=8192), + ) + act = MoEActivationPack(x_q, x_sf, topk_ids, topk_weights) + weights = MoEWeightPack() + weights.prepare_for("cutlass_mxfp8", view) + runner = MoELayer(config).runners[0] + inputs = runner.pack_inputs(act, weights) + assert x_sf.numel() == _mxfp8_swizzled_act_sf_numel(num_tokens, hidden_size) + assert inputs[-1].numel() == x_sf.numel() + bucket = map_to_hybrid_bucket(num_tokens, 8192) + assert bucket == 512 + assert runner.tuning_config.constraint_specs + infer_numel = runner.tuning_config.constraint_specs[0].infer_shape + assert infer_numel([None, (bucket, hidden_size)]) == _mxfp8_swizzled_act_sf_numel( + bucket, hidden_size + ) + + synthesized = list(inputs) + synthesized[0] = torch.empty( + bucket, hidden_size, dtype=torch.bfloat16, device=device + ) + synthesized[1] = torch.empty( + bucket, hidden_size, dtype=torch.float8_e4m3fn, device=device + ) + synthesized[2] = torch.empty(bucket, top_k, dtype=torch.int32, device=device) + synthesized[3] = torch.empty(bucket, top_k, dtype=torch.float32, device=device) + tuned = runner._prepare_tuning_inputs(synthesized) + assert tuned[-1].numel() == _mxfp8_swizzled_act_sf_numel(bucket, hidden_size) + + with autotune(True): + _, tactic = AutoTuner.get().choose_one( + "test_moe_cutlass_mxfp8_bucket_boundary", + [runner], + runner.tuning_config, + inputs, + ) + assert isinstance(tactic, tuple) and len(tactic) == 2 + assert all(stage_tactic >= 0 for stage_tactic in tactic) + actual = runner.forward(inputs, tactic=tactic) + torch.cuda.synchronize() + assert torch.isfinite(actual).all() + + +def _dequant_int4(packed, scale, group_size=128): + even = packed.to(torch.int16) & 0xF + odd = packed.to(torch.int16) >> 4 + even = torch.where(even >= 8, even - 16, even) + odd = torch.where(odd >= 8, odd - 16, odd) + unpacked = torch.stack((even, odd), dim=-1).reshape( + *packed.shape[:-1], packed.shape[-1] * 2 + ) + expanded = scale.float().repeat_interleave(group_size, dim=-1) + return unpacked.float() * expanded + + +@cutlass_w4a8_required +def test_cutlass_w4a8_moe_layer_matches_quantized_reference(): + torch.manual_seed(49) + device = torch.device("cuda", torch.cuda.current_device()) + num_tokens, num_experts, top_k = 16, 4, 2 + hidden_size, intermediate_size = 128, 256 + x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) / 2 + w1, w2 = _make_bf16_experts(num_experts, hidden_size, intermediate_size, device) + topk_ids, topk_weights = _make_routing(num_tokens, num_experts, top_k, device) + from flashinfer.fused_moe.prepare import _quantize_int4_grouped + + packed_w1, scale_w1 = _quantize_int4_grouped(w1) + packed_w2, scale_w2 = _quantize_int4_grouped(w2) + view = CutlassW4A8Config.prepare_weights( + w1, + w2, + num_local_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + config = _config( + quant=QuantConfig(variant=QuantVariant.W4A8), + experts=ExpertConfig(intermediate_size=intermediate_size), + backend=BackendOptions((CutlassW4A8Config(),)), + execution=ExecutionConfig(enable_pdl=False, tune_max_num_tokens=64), + ) + act = MoEActivationPack(x, None, topk_ids, topk_weights) + weights = MoEWeightPack() + weights.prepare_for("cutlass_w4a8", view) + layer = MoELayer(config) + _pin_fallback_winner(layer, act) + actual = layer(act, weights) + expected = _reference( + act, + _dequant_int4(packed_w1, scale_w1).to(torch.bfloat16), + _dequant_int4(packed_w2, scale_w2).to(torch.bfloat16), + ) + assert layer.winner_backend == "cutlass_w4a8" + _assert_numerically_close(actual, expected, rtol=1e-1, atol=1e-1) + _autotune_and_graph( + layer.runners[0], + act, + weights, + expected, + rtol=1e-1, + atol=1e-1, + cache_name="test_moe_cutlass_w4a8", + ) + + +@cutlass_humming_required +def test_cutlass_humming_moe_layer_matches_quantized_reference(): + torch.manual_seed(50) + device = torch.device("cuda", torch.cuda.current_device()) + num_tokens, num_experts, top_k = 16, 4, 2 + hidden_size, intermediate_size = 128, 256 + x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) / 2 + w1, w2 = _make_bf16_experts(num_experts, hidden_size, intermediate_size, device) + topk_ids, topk_weights = _make_routing(num_tokens, num_experts, top_k, device) + view = CutlassHummingConfig.prepare_weights( + w1, + w2, + num_local_experts=num_experts, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + device=device, + ) + w1_lin, w1_sf = _quantize_mxfp4_linear( + w1.view(num_experts * 2 * intermediate_size, hidden_size) + ) + w2_lin, w2_sf = _quantize_mxfp4_linear( + w2.view(num_experts * hidden_size, intermediate_size) + ) + config = _config( + quant=QuantConfig(variant=QuantVariant.Humming), + experts=ExpertConfig(intermediate_size=intermediate_size), + backend=BackendOptions((CutlassHummingConfig(),)), + execution=ExecutionConfig(enable_pdl=False, tune_max_num_tokens=64), + ) + act = MoEActivationPack(x, None, topk_ids, topk_weights) + weights = MoEWeightPack() + weights.prepare_for("cutlass_humming", view) + layer = MoELayer(config) + _pin_fallback_winner(layer, act) + actual = layer(act, weights) + expected = _reference( + act, + _dequant_linear_mxfp4( + w1_lin.view(num_experts, 2 * intermediate_size, hidden_size // 2), + w1_sf.view(num_experts, 2 * intermediate_size, hidden_size // 32), + ).to(torch.bfloat16), + _dequant_linear_mxfp4( + w2_lin.view(num_experts, hidden_size, intermediate_size // 2), + w2_sf.view(num_experts, hidden_size, intermediate_size // 32), + ).to(torch.bfloat16), + ) + assert layer.winner_backend == "cutlass_humming" + _assert_numerically_close(actual, expected, rtol=2e-1, atol=2e-1) + _autotune_and_graph( + layer.runners[0], + act, + weights, + expected, + rtol=2e-1, + atol=2e-1, + cache_name="test_moe_cutlass_humming", + )