Implement SM120 DeepSeek V4 flashinfer_mxfp4 moe runner backend + TP2 (#30272)
This commit is contained in:
@@ -149,7 +149,7 @@ import { Playground } from "/src/snippets/_playground.jsx";
|
|||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><strong><a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash">DeepSeek-V4-Flash</a></strong></td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><strong><a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Flash">DeepSeek-V4-Flash</a></strong></td>
|
||||||
<td style={{padding: "9px 12px", textAlign: "right", backgroundColor: "rgba(255,255,255,0.05)"}}><strong>284B</strong></td>
|
<td style={{padding: "9px 12px", textAlign: "right", backgroundColor: "rgba(255,255,255,0.05)"}}><strong>284B</strong></td>
|
||||||
<td style={{padding: "9px 12px", textAlign: "right", backgroundColor: "rgba(255,255,255,0.02)"}}>13B</td>
|
<td style={{padding: "9px 12px", textAlign: "right", backgroundColor: "rgba(255,255,255,0.02)"}}>13B</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>single-node serving on B200 / B300 / GB200 / GB300 / H200 (TP=4); H100 (TP=8)</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>single-node serving on B200 / B300 / GB200 / GB300 / H200 (TP=4); RTX PRO 6000 (TP=2); H100 (TP=8)</td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><strong><a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro">DeepSeek-V4-Pro</a></strong></td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><strong><a href="https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro">DeepSeek-V4-Pro</a></strong></td>
|
||||||
@@ -293,9 +293,8 @@ TCP, which can lead to garbled KV transfer on large checkpoints.
|
|||||||
|
|
||||||
**RTX PRO 6000 (SM120 / Blackwell Desktop) note**
|
**RTX PRO 6000 (SM120 / Blackwell Desktop) note**
|
||||||
|
|
||||||
RTX PRO 6000 (96 GB) runs **Flash only** — V4-Pro doesn't fit on 8× 96 GB. It uses the
|
RTX PRO 6000 (96 GB) runs **Flash only** with the FlashInfer MXFP4 MoE runner.
|
||||||
**low-latency / TP-only** recipe (TP=4, single node) with the **Marlin** W4A16 MoE runner and
|
V4-Pro doesn't fit on 8× 96 GB; the Deploy panel greys out unsupported recipes.
|
||||||
`--mem-fraction-static 0.70`; the Deploy panel greys out the other recipes for this card.
|
|
||||||
HiCache and MegaMoE are **not** supported on RTX PRO 6000.
|
HiCache and MegaMoE are **not** supported on RTX PRO 6000.
|
||||||
|
|
||||||
**AMD (MI300X / MI355X) note**
|
**AMD (MI300X / MI355X) note**
|
||||||
|
|||||||
@@ -1421,7 +1421,6 @@ sgl-eval run aime25 \\
|
|||||||
|
|
||||||
// ====================================================================
|
// ====================================================================
|
||||||
// RTX PRO 6000 (SM120 / Blackwell Desktop) — Flash + low-latency only
|
// RTX PRO 6000 (SM120 / Blackwell Desktop) — Flash + low-latency only
|
||||||
// (V4-Pro doesn't fit on 8× 96 GB); TP-only, Marlin MoE runner.
|
|
||||||
// ====================================================================
|
// ====================================================================
|
||||||
{
|
{
|
||||||
match: { hw: "rtx6000", variant: "flash", quant: "fp4", strategy: "low-latency", nodes: "single" },
|
match: { hw: "rtx6000", variant: "flash", quant: "fp4", strategy: "low-latency", nodes: "single" },
|
||||||
@@ -1430,9 +1429,9 @@ sgl-eval run aime25 \\
|
|||||||
flags: [
|
flags: [
|
||||||
"--trust-remote-code",
|
"--trust-remote-code",
|
||||||
"--model-path {{MODEL_NAME}}",
|
"--model-path {{MODEL_NAME}}",
|
||||||
"--tp 4",
|
"--tp 2",
|
||||||
"--moe-runner-backend marlin",
|
"--moe-runner-backend flashinfer_mxfp4",
|
||||||
"--mem-fraction-static 0.70",
|
"--mem-fraction-static 0.92",
|
||||||
"--cuda-graph-max-bs-decode 32",
|
"--cuda-graph-max-bs-decode 32",
|
||||||
"--host {{HOST_IP}}",
|
"--host {{HOST_IP}}",
|
||||||
"--port {{PORT}}",
|
"--port {{PORT}}",
|
||||||
|
|||||||
@@ -375,14 +375,11 @@ __global__ void Marlin(
|
|||||||
is_zp_float ? prob_n * prob_k / group_size / 8 : prob_n * prob_k / group_size / (pack_factor * 4);
|
is_zp_float ? prob_n * prob_k / group_size / 8 : prob_n * prob_k / group_size / (pack_factor * 4);
|
||||||
const int b_bias_expert_stride = prob_n / 8;
|
const int b_bias_expert_stride = prob_n / 8;
|
||||||
|
|
||||||
// parallel: num valid moe blocks
|
|
||||||
int num_tokens_past_padded = num_tokens_past_padded_ptr[0];
|
int num_tokens_past_padded = num_tokens_past_padded_ptr[0];
|
||||||
int parallel = num_tokens_past_padded / moe_block_size;
|
int parallel = num_tokens_past_padded / moe_block_size;
|
||||||
int num_valid_blocks = parallel;
|
int num_valid_blocks = parallel;
|
||||||
if (is_ep) {
|
for (int i = 0; i < parallel; i++) {
|
||||||
for (int i = 0; i < parallel; i++) {
|
if (expert_ids_ptr[i] == -1) num_valid_blocks--;
|
||||||
if (expert_ids_ptr[i] == -1) num_valid_blocks--;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
int num_invalid_blocks = parallel - num_valid_blocks;
|
int num_invalid_blocks = parallel - num_valid_blocks;
|
||||||
parallel = num_valid_blocks;
|
parallel = num_valid_blocks;
|
||||||
|
|||||||
@@ -1443,11 +1443,14 @@ def fp8_paged_mqa_logits_kernel(
|
|||||||
for j in T.Pipelined(n_iters, num_stages=2):
|
for j in T.Pipelined(n_iters, num_stages=2):
|
||||||
i = i_start + j
|
i = i_start + j
|
||||||
page = page_table[bx, i]
|
page = page_table[bx, i]
|
||||||
k_smem_u8 = T.alloc_shared((B * D,), UINT8)
|
k_smem_u8 = T.alloc_shared((1, B * D), UINT8)
|
||||||
T.copy(kvcache_u8[page, 0:SCALE_OFFSET], k_smem_u8)
|
T.copy(kvcache_u8[page : page + 1, 0:SCALE_OFFSET], k_smem_u8)
|
||||||
k_smem = T.view(k_smem_u8, (B, D), FP8)
|
k_smem = T.view(k_smem_u8, (B, D), FP8)
|
||||||
k_s_smem_u8 = T.alloc_shared((B * 4,), UINT8)
|
k_s_smem_u8 = T.alloc_shared((1, B * 4), UINT8)
|
||||||
T.copy(kvcache_u8[page, SCALE_OFFSET:BLOCK_BYTES], k_s_smem_u8)
|
T.copy(
|
||||||
|
kvcache_u8[page : page + 1, SCALE_OFFSET:BLOCK_BYTES],
|
||||||
|
k_s_smem_u8,
|
||||||
|
)
|
||||||
k_s_smem = T.view(k_s_smem_u8, (B,), FP32)
|
k_s_smem = T.view(k_s_smem_u8, (B,), FP32)
|
||||||
k_s_frag = T.alloc_fragment((B,), FP32)
|
k_s_frag = T.alloc_fragment((B,), FP32)
|
||||||
T.copy(k_s_smem, k_s_frag)
|
T.copy(k_s_smem, k_s_frag)
|
||||||
|
|||||||
@@ -400,14 +400,19 @@ def _flash_mla_flashinfer(
|
|||||||
extra_indices,
|
extra_indices,
|
||||||
extra_topk_length,
|
extra_topk_length,
|
||||||
):
|
):
|
||||||
"""FlashInfer SM120 sparse MLA via sparse_mla_sm120_decode_dsv4.
|
"""FlashInfer SM120 sparse MLA via the paged-attention dispatcher.
|
||||||
|
|
||||||
SGLang SWA pool uses page_size=256 (footer format: 256*576 bytes data + 256*8 bytes scale).
|
SGLang SWA pool uses page_size=256 (footer format: 256*576 bytes data + 256*8 bytes scale).
|
||||||
FlashInfer decode_dsv4 fast path requires page_block_size=64 (footer: 64*576 + 64*8).
|
FlashInfer decode_dsv4 fast path requires page_block_size=64 (footer: 64*576 + 64*8).
|
||||||
We split 256-token pages into 4 virtual 64-token pages.
|
We split 256-token pages into 4 virtual 64-token pages.
|
||||||
Token indices are invariant under page-split (identity mapping).
|
Token indices are invariant under page-split (identity mapping).
|
||||||
"""
|
"""
|
||||||
from flashinfer.mla._sparse_mla_sm120 import sparse_mla_sm120_decode_dsv4
|
from flashinfer.mla._sparse_mla_sm120 import (
|
||||||
|
_DECODE_MAX_TOKENS as _FI_DECODE_MAX_TOKENS,
|
||||||
|
)
|
||||||
|
from flashinfer.mla._sparse_mla_sm120 import (
|
||||||
|
_sparse_mla_sm120_paged_attention,
|
||||||
|
)
|
||||||
|
|
||||||
B, _, H, D = q.shape # (batch, 1, num_heads, head_dim)
|
B, _, H, D = q.shape # (batch, 1, num_heads, head_dim)
|
||||||
dev = q.device
|
dev = q.device
|
||||||
@@ -435,32 +440,37 @@ def _flash_mla_flashinfer(
|
|||||||
output = torch.empty(B, H, head_dim_v, dtype=torch.bfloat16, device=dev)
|
output = torch.empty(B, H, head_dim_v, dtype=torch.bfloat16, device=dev)
|
||||||
out_lse = torch.empty(B, H, dtype=torch.float32, device=dev)
|
out_lse = torch.empty(B, H, dtype=torch.float32, device=dev)
|
||||||
|
|
||||||
# Pre-allocate split-K scratch for decode-dsv4 fast path.
|
# Use split-K for decode-sized batches and paged attention otherwise.
|
||||||
topk = idx.shape[-1]
|
if B <= _FI_DECODE_MAX_TOKENS:
|
||||||
extra_topk = extra_idx.shape[-1] if extra_idx is not None else 0
|
topk = idx.shape[-1]
|
||||||
_BI = 64
|
extra_topk = extra_idx.shape[-1] if extra_idx is not None else 0
|
||||||
num_splits = (topk + _BI - 1) // _BI + (
|
_BI = 64
|
||||||
(extra_topk + _BI - 1) // _BI if extra_topk > 0 else 0
|
num_splits = (topk + _BI - 1) // _BI + (
|
||||||
)
|
(extra_topk + _BI - 1) // _BI if extra_topk > 0 else 0
|
||||||
mid_out = torch.empty(
|
)
|
||||||
B, H, num_splits, head_dim_v, dtype=torch.bfloat16, device=dev
|
mid_out = torch.empty(
|
||||||
)
|
B, H, num_splits, head_dim_v, dtype=torch.bfloat16, device=dev
|
||||||
mid_lse = torch.empty(B, H, num_splits, dtype=torch.float32, device=dev)
|
)
|
||||||
|
mid_lse = torch.empty(B, H, num_splits, dtype=torch.float32, device=dev)
|
||||||
|
else:
|
||||||
|
mid_out = None
|
||||||
|
mid_lse = None
|
||||||
|
|
||||||
sparse_mla_sm120_decode_dsv4(
|
_sparse_mla_sm120_paged_attention(
|
||||||
q=q.squeeze(1) if q.ndim == 4 else q,
|
q.squeeze(1) if q.ndim == 4 else q,
|
||||||
kv_cache=kv_64,
|
kv_64,
|
||||||
indices=idx,
|
idx,
|
||||||
mid_out=mid_out,
|
output,
|
||||||
mid_lse=mid_lse,
|
out_lse,
|
||||||
output=output,
|
softmax_scale,
|
||||||
out_lse=out_lse,
|
d_v=head_dim_v,
|
||||||
sm_scale=softmax_scale,
|
|
||||||
topk_length=topk_length,
|
topk_length=topk_length,
|
||||||
attn_sink=attn_sink,
|
attn_sink=attn_sink,
|
||||||
extra_kv_cache=extra_kv_64,
|
extra_kv_cache=extra_kv_64,
|
||||||
extra_indices=extra_idx,
|
extra_indices=extra_idx,
|
||||||
extra_topk_length=extra_topk_length,
|
extra_topk_length=extra_topk_length,
|
||||||
|
mid_out=mid_out,
|
||||||
|
mid_lse=mid_lse,
|
||||||
)
|
)
|
||||||
|
|
||||||
return (output.unsqueeze(1), None)
|
return (output.unsqueeze(1), None)
|
||||||
|
|||||||
@@ -1434,16 +1434,15 @@ def _deepseek_v4_kv_cache_dtype(view: Any) -> dict:
|
|||||||
|
|
||||||
@register_post_process
|
@register_post_process
|
||||||
def _deepseek_v4_sm120_moe(view: Any) -> dict:
|
def _deepseek_v4_sm120_moe(view: Any) -> dict:
|
||||||
"""Slot pass in the DeepSeek V4 validation branch: SM120 lacks
|
"""Default DeepSeek V4 MXFP4 experts to FlashInfer CUTLASS on SM120."""
|
||||||
tcgen05/TMEM, fall back to the marlin MoE runner (reads the
|
|
||||||
mid-resolution moe_runner_backend, after the dispatch-time nvfp4
|
|
||||||
default)."""
|
|
||||||
hf_config = view.get_model_config().hf_config
|
hf_config = view.get_model_config().hf_config
|
||||||
if hf_config.architectures[0] != "DeepseekV4ForCausalLM":
|
if hf_config.architectures[0] != "DeepseekV4ForCausalLM":
|
||||||
return {}
|
return {}
|
||||||
if is_sm120_supported() and view.moe_runner_backend == "auto":
|
if is_sm120_supported() and view.moe_runner_backend == "auto":
|
||||||
logger.info("Use marlin as MoE runner backend on SM120 for DeepseekV4")
|
logger.info(
|
||||||
return {"moe_runner_backend": "marlin"}
|
"Use flashinfer_mxfp4 as MoE runner backend on SM120 for DeepseekV4"
|
||||||
|
)
|
||||||
|
return {"moe_runner_backend": "flashinfer_mxfp4"}
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1688,9 +1688,14 @@ class DeepseekV4AttnBackend(
|
|||||||
extra_indices.shape[-1] % 64 == 0
|
extra_indices.shape[-1] % 64 == 0
|
||||||
), f"{extra_indices.shape=}'s last dimension is not aligned to 64"
|
), f"{extra_indices.shape=}'s last dimension is not aligned to 64"
|
||||||
|
|
||||||
if forward_batch.forward_mode.is_extend_without_speculative() and (
|
# sparse_prefill_fwd does not support SM120.
|
||||||
q.shape[0] > _LARGE_INDEXER_QUERY_THRESHOLD
|
if (
|
||||||
or envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.get()
|
forward_batch.forward_mode.is_extend_without_speculative()
|
||||||
|
and not _is_sm120
|
||||||
|
and (
|
||||||
|
q.shape[0] > _LARGE_INDEXER_QUERY_THRESHOLD
|
||||||
|
or envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.get()
|
||||||
|
)
|
||||||
):
|
):
|
||||||
return self._forward_prefill_sparse(
|
return self._forward_prefill_sparse(
|
||||||
q=q,
|
q=q,
|
||||||
|
|||||||
@@ -174,6 +174,25 @@ def fp8_paged_mqa_logits_torch_sm120(
|
|||||||
block_size = kvcache_fp8.shape[1]
|
block_size = kvcache_fp8.shape[1]
|
||||||
device = q_fp8.device
|
device = q_fp8.device
|
||||||
|
|
||||||
|
_QUERY_CHUNK = 1024
|
||||||
|
if batch_size > _QUERY_CHUNK:
|
||||||
|
return torch.cat(
|
||||||
|
[
|
||||||
|
fp8_paged_mqa_logits_torch_sm120(
|
||||||
|
q_fp8[start : start + _QUERY_CHUNK],
|
||||||
|
kvcache_fp8,
|
||||||
|
weight[start : start + _QUERY_CHUNK],
|
||||||
|
seq_lens[start : start + _QUERY_CHUNK],
|
||||||
|
page_table[start : start + _QUERY_CHUNK],
|
||||||
|
deep_gemm_metadata,
|
||||||
|
max_seq_len,
|
||||||
|
clean_logits=clean_logits,
|
||||||
|
)
|
||||||
|
for start in range(0, batch_size, _QUERY_CHUNK)
|
||||||
|
],
|
||||||
|
dim=0,
|
||||||
|
)
|
||||||
|
|
||||||
assert head_dim == 128, "Vectorized torch impl hardcodes DSV4 indexer head_dim=128"
|
assert head_dim == 128, "Vectorized torch impl hardcodes DSV4 indexer head_dim=128"
|
||||||
assert (
|
assert (
|
||||||
block_size == 64
|
block_size == 64
|
||||||
|
|||||||
@@ -204,7 +204,8 @@ def fused_marlin_moe(
|
|||||||
device=hidden_states.device,
|
device=hidden_states.device,
|
||||||
dtype=hidden_states.dtype,
|
dtype=hidden_states.dtype,
|
||||||
)
|
)
|
||||||
intermediate_cache13 = torch.empty(
|
# Marlin skips masked expert rows, so their shared cache must start at zero.
|
||||||
|
intermediate_cache13 = torch.zeros(
|
||||||
(M * topk_ids.shape[1] * max(gemm1_n, K),),
|
(M * topk_ids.shape[1] * max(gemm1_n, K),),
|
||||||
device=hidden_states.device,
|
device=hidden_states.device,
|
||||||
dtype=hidden_states.dtype,
|
dtype=hidden_states.dtype,
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
"""FlashInfer CUTLASS MoE fused funcs.
|
"""FlashInfer CUTLASS MoE fused funcs.
|
||||||
|
|
||||||
This module owns the FlashInfer ``cutlass_fused_moe`` calls used by the
|
This module owns the FlashInfer ``cutlass_fused_moe`` calls used by the
|
||||||
unquantized, ModelOpt FP8, ModelOpt NVFP4, and SM90 MXFP4 MoE paths.
|
unquantized, ModelOpt FP8, ModelOpt NVFP4, and MXFP4 MoE paths.
|
||||||
Quantization methods prepare a small quant_info payload and route through
|
Quantization methods prepare a small quant_info payload and route through
|
||||||
``MoeRunner``.
|
``MoeRunner``.
|
||||||
"""
|
"""
|
||||||
@@ -62,27 +62,28 @@ class FlashInferCutlassMoeQuantInfo(MoeQuantInfo):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class FlashInferCutlassMxfp4MoeQuantInfo(MoeQuantInfo):
|
class FlashInferCutlassMxfp4MoeQuantInfo(MoeQuantInfo):
|
||||||
"""Quantization payload for the SM90 CUTLASS W4A16 MXFP4 MoE path.
|
"""Quantization payload for CUTLASS MXFP4 MoE.
|
||||||
|
|
||||||
Weights and scales are pre-interleaved at load time via
|
SM90 consumes W4A16-interleaved weights and scales. SM120 consumes packed
|
||||||
``interleave_moe_{weights,scales}_for_sm90_mixed_gemm``; this dataclass
|
MXFP4 weights and block-interleaved scales with MXFP8 activations.
|
||||||
only carries references plus the per-call routing/topology fields.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Pre-interleaved weights (uint8, packed FP4)
|
# SM90 weights are interleaved; SM120 weights remain checkpoint-packed.
|
||||||
w13_weight: torch.Tensor # [E, 2*N, K/2]
|
w13_weight: torch.Tensor # [E, 2*N, K/2]
|
||||||
w2_weight: torch.Tensor # [E, K, N/2]
|
w2_weight: torch.Tensor # [E, K, N/2]
|
||||||
|
|
||||||
# Pre-interleaved E8M0 block scales (uint8; viewed as int32 at call time)
|
# E8M0 block scales in the layout selected by the quantization method.
|
||||||
w13_weight_scale: torch.Tensor # [E, 2*N, K/32]
|
w13_weight_scale: torch.Tensor # [E, 2*N, K/32]
|
||||||
w2_weight_scale: torch.Tensor # [E, K, N/32]
|
w2_weight_scale: torch.Tensor # [E, K, N/32]
|
||||||
|
|
||||||
|
# A non-None global scale selects the SM120 MXFP8 activation path.
|
||||||
|
mxfp4_weight_global_scale: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# Per-expert bias. GPT-OSS has both; DSv4 leaves both None.
|
# Per-expert bias. GPT-OSS has both; DSv4 leaves both None.
|
||||||
w13_bias: Optional[torch.Tensor] = None # bf16 [E, 2*N]
|
w13_bias: Optional[torch.Tensor] = None # bf16 [E, 2*N]
|
||||||
w2_bias: Optional[torch.Tensor] = None # bf16 [E, K]
|
w2_bias: Optional[torch.Tensor] = None # bf16 [E, K]
|
||||||
|
|
||||||
# Per-expert SwiGLU scalars (fp32 [E]). Either all three are present
|
# Optional per-expert SwiGLU overrides, fp32 [E].
|
||||||
# (clamped SwiGLU) or all three are None (kernel default SwiGLU).
|
|
||||||
swiglu_alpha: Optional[torch.Tensor] = None
|
swiglu_alpha: Optional[torch.Tensor] = None
|
||||||
swiglu_beta: Optional[torch.Tensor] = None
|
swiglu_beta: Optional[torch.Tensor] = None
|
||||||
swiglu_limit: Optional[torch.Tensor] = None
|
swiglu_limit: Optional[torch.Tensor] = None
|
||||||
@@ -297,11 +298,7 @@ def fused_experts_none_to_flashinfer_mxfp4(
|
|||||||
quant_info: MoeQuantInfo,
|
quant_info: MoeQuantInfo,
|
||||||
runner_config: MoeRunnerConfig,
|
runner_config: MoeRunnerConfig,
|
||||||
) -> StandardCombineInput:
|
) -> StandardCombineInput:
|
||||||
"""SM90 W4A16 MXFP4 fused expert forward pass.
|
"""Run the FlashInfer CUTLASS MXFP4 fused experts."""
|
||||||
|
|
||||||
This preserves the ``flashinfer_mxfp4`` runner backend registration while
|
|
||||||
centralizing the CUTLASS execution in this module.
|
|
||||||
"""
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
||||||
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
||||||
|
|
||||||
@@ -335,6 +332,33 @@ def fused_experts_none_to_flashinfer_mxfp4(
|
|||||||
value=0.0,
|
value=0.0,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
weight_global_scale = quant_info.mxfp4_weight_global_scale
|
||||||
|
use_mxfp8_act_scaling = weight_global_scale is not None
|
||||||
|
input_sf = None
|
||||||
|
fc1_expert_weights = quant_info.w13_weight
|
||||||
|
fc2_expert_weights = quant_info.w2_weight
|
||||||
|
if weight_global_scale is not None:
|
||||||
|
from flashinfer import mxfp8_quantize
|
||||||
|
|
||||||
|
x, input_sf = mxfp8_quantize(
|
||||||
|
x,
|
||||||
|
is_sf_swizzled_layout=True,
|
||||||
|
alignment=32,
|
||||||
|
)
|
||||||
|
fc1_expert_weights = fc1_expert_weights.view(torch.int64)
|
||||||
|
fc2_expert_weights = fc2_expert_weights.view(torch.int64)
|
||||||
|
quant_scales = [
|
||||||
|
quant_info.w13_weight_scale.view(torch.int32),
|
||||||
|
weight_global_scale,
|
||||||
|
quant_info.w2_weight_scale.view(torch.int32),
|
||||||
|
weight_global_scale,
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
quant_scales = [
|
||||||
|
quant_info.w13_weight_scale.view(torch.int32),
|
||||||
|
quant_info.w2_weight_scale.view(torch.int32),
|
||||||
|
]
|
||||||
|
|
||||||
out_hidden = padded_hidden if do_pad else origin_hidden
|
out_hidden = padded_hidden if do_pad else origin_hidden
|
||||||
output_dtype = torch.bfloat16
|
output_dtype = torch.bfloat16
|
||||||
with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()):
|
with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()):
|
||||||
@@ -342,15 +366,13 @@ def fused_experts_none_to_flashinfer_mxfp4(
|
|||||||
|
|
||||||
flashinfer_cutlass_fused_moe(
|
flashinfer_cutlass_fused_moe(
|
||||||
input=x,
|
input=x,
|
||||||
token_selected_experts=topk_ids.to(torch.int),
|
token_selected_experts=topk_ids.to(torch.int32),
|
||||||
token_final_scales=topk_weights,
|
token_final_scales=topk_weights,
|
||||||
fc1_expert_weights=quant_info.w13_weight,
|
fc1_expert_weights=fc1_expert_weights,
|
||||||
fc2_expert_weights=quant_info.w2_weight,
|
fc2_expert_weights=fc2_expert_weights,
|
||||||
output_dtype=output_dtype,
|
output_dtype=output_dtype,
|
||||||
quant_scales=[
|
quant_scales=quant_scales,
|
||||||
quant_info.w13_weight_scale.view(torch.int32),
|
input_sf=input_sf,
|
||||||
quant_info.w2_weight_scale.view(torch.int32),
|
|
||||||
],
|
|
||||||
fc1_expert_biases=quant_info.w13_bias,
|
fc1_expert_biases=quant_info.w13_bias,
|
||||||
fc2_expert_biases=quant_info.w2_bias,
|
fc2_expert_biases=quant_info.w2_bias,
|
||||||
swiglu_alpha=quant_info.swiglu_alpha,
|
swiglu_alpha=quant_info.swiglu_alpha,
|
||||||
@@ -360,7 +382,8 @@ def fused_experts_none_to_flashinfer_mxfp4(
|
|||||||
tp_rank=quant_info.moe_tp_rank,
|
tp_rank=quant_info.moe_tp_rank,
|
||||||
ep_size=quant_info.moe_ep_size,
|
ep_size=quant_info.moe_ep_size,
|
||||||
ep_rank=quant_info.moe_ep_rank,
|
ep_rank=quant_info.moe_ep_rank,
|
||||||
use_w4_group_scaling=True,
|
use_w4_group_scaling=not use_mxfp8_act_scaling,
|
||||||
|
use_mxfp8_act_scaling=use_mxfp8_act_scaling,
|
||||||
activation_type=ActivationType.Swiglu,
|
activation_type=ActivationType.Swiglu,
|
||||||
tune_max_num_tokens=next_power_of_2(x.shape[0]),
|
tune_max_num_tokens=next_power_of_2(x.shape[0]),
|
||||||
output=out,
|
output=out,
|
||||||
|
|||||||
@@ -383,9 +383,8 @@ class Fp8Config(QuantizationConfig):
|
|||||||
return Mxfp4HummingMoEMethod(fp8_method, prefix=prefix)
|
return Mxfp4HummingMoEMethod(fp8_method, prefix=prefix)
|
||||||
|
|
||||||
if self.is_fp4_experts and get_moe_runner_backend().is_flashinfer_mxfp4():
|
if self.is_fp4_experts and get_moe_runner_backend().is_flashinfer_mxfp4():
|
||||||
# SM100 (Blackwell) -> trtllm-gen path.
|
# SM100 uses TRT-LLM; SM90 uses W4A16 and SM120 uses MXFP8xMXFP4.
|
||||||
# SM90 (Hopper) -> cutlass mixed-input path (FlashInfer #3084).
|
if is_sm90_supported() or is_sm120_supported():
|
||||||
if is_sm90_supported() and not is_sm100_supported():
|
|
||||||
from sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe import (
|
from sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe import (
|
||||||
Mxfp4FlashinferCutlassMoEMethod,
|
Mxfp4FlashinferCutlassMoEMethod,
|
||||||
)
|
)
|
||||||
@@ -1053,6 +1052,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
intermediate_size_per_partition: int,
|
intermediate_size_per_partition: int,
|
||||||
params_dtype: torch.dtype,
|
params_dtype: torch.dtype,
|
||||||
with_bias: bool = False,
|
with_bias: bool = False,
|
||||||
|
fp4_scale_dtype: Optional[torch.dtype] = None,
|
||||||
**extra_weight_attrs,
|
**extra_weight_attrs,
|
||||||
):
|
):
|
||||||
self.with_bias = with_bias
|
self.with_bias = with_bias
|
||||||
@@ -1188,7 +1188,8 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
# WEIGHT_SCALES
|
# WEIGHT_SCALES
|
||||||
if self.is_fp4_expert:
|
if self.is_fp4_expert:
|
||||||
fp4_block_k = 32
|
fp4_block_k = 32
|
||||||
fp4_scale_dtype = torch.float8_e8m0fnu if _use_aiter else torch.float32
|
if fp4_scale_dtype is None:
|
||||||
|
fp4_scale_dtype = torch.float8_e8m0fnu if _use_aiter else torch.float32
|
||||||
w13_weight_scale = torch.nn.Parameter(
|
w13_weight_scale = torch.nn.Parameter(
|
||||||
torch.ones(
|
torch.ones(
|
||||||
num_experts,
|
num_experts,
|
||||||
|
|||||||
@@ -399,6 +399,11 @@ def prepare_moe_mxfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
|
|||||||
_permute_bias(w2_bias_data), requires_grad=False
|
_permute_bias(w2_bias_data), requires_grad=False
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Marlin uses the repacked scales; release the loader-format parameters.
|
||||||
|
for stale in ("w13_weight_scale_inv", "w2_weight_scale_inv"):
|
||||||
|
if hasattr(layer, stale):
|
||||||
|
delattr(layer, stale)
|
||||||
|
|
||||||
|
|
||||||
def prepare_moe_nvfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
|
def prepare_moe_nvfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
|
||||||
if layer.quant_config.group_size != 16:
|
if layer.quant_config.group_size != 16:
|
||||||
|
|||||||
@@ -1,19 +1,7 @@
|
|||||||
"""DeepSeek-V4 MXFP4 expert backend backed by FlashInfer's SM90 cutlass
|
"""DeepSeek-V4 MXFP4 expert backend backed by FlashInfer CUTLASS MoE.
|
||||||
mixed-input MoE GEMM (FlashInfer PR #3084).
|
|
||||||
|
|
||||||
Sibling of :class:`Mxfp4MarlinMoEMethod` and :class:`Mxfp4FlashinferTrtllmMoEMethod`.
|
``Fp8Config`` selects this backend for SM90 and SM120; SM100 uses the
|
||||||
Wired into :func:`Fp8MoEConfig.get_quant_method` when
|
TRT-LLM implementation.
|
||||||
``is_fp4_experts=True`` and ``--moe-runner-backend flashinfer_mxfp4`` is
|
|
||||||
selected on a Hopper (SM90) device. SM100 still routes to
|
|
||||||
:class:`Mxfp4FlashinferTrtllmMoEMethod` (trtllm-gen).
|
|
||||||
|
|
||||||
Performance trade-off vs Marlin (kernel-level on H100, GPT-OSS-like body):
|
|
||||||
- decode (M <= 64) : Marlin +12-15 %
|
|
||||||
- tie (M ~= 256)
|
|
||||||
- prefill (M >= 1024) : FlashInfer +24-36 %
|
|
||||||
|
|
||||||
PD-disaggregated prefill workers are the natural fit; decode workers should
|
|
||||||
keep the Marlin default.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -26,29 +14,12 @@ import torch
|
|||||||
from torch.nn import Module
|
from torch.nn import Module
|
||||||
from torch.nn.parameter import Parameter
|
from torch.nn.parameter import Parameter
|
||||||
|
|
||||||
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
|
||||||
from sglang.srt.utils import is_flashinfer_available, log_info_on_rank0
|
from sglang.srt.utils import is_flashinfer_available, log_info_on_rank0
|
||||||
|
from sglang.srt.utils.common import is_sm120_supported
|
||||||
|
|
||||||
# Silence the TRT-LLM cutlass autotune trace embedded inside FlashInfer's
|
# Suppress TRT-LLM CUTLASS trace logs without overriding user configuration.
|
||||||
# cutlass_fused_moe. Its C++ logger reads TLLM_LOG_LEVEL on first kernel launch;
|
|
||||||
# setdefault preserves any explicit user override.
|
|
||||||
os.environ.setdefault("TLLM_LOG_LEVEL", "INFO")
|
os.environ.setdefault("TLLM_LOG_LEVEL", "INFO")
|
||||||
|
|
||||||
if is_flashinfer_available():
|
|
||||||
try:
|
|
||||||
from flashinfer.fused_moe import (
|
|
||||||
interleave_moe_scales_for_sm90_mixed_gemm,
|
|
||||||
interleave_moe_weights_for_sm90_mixed_gemm,
|
|
||||||
)
|
|
||||||
|
|
||||||
_FI_HAS_SM90_CUTLASS_MXFP4 = True
|
|
||||||
except ImportError:
|
|
||||||
interleave_moe_scales_for_sm90_mixed_gemm = None
|
|
||||||
interleave_moe_weights_for_sm90_mixed_gemm = None
|
|
||||||
_FI_HAS_SM90_CUTLASS_MXFP4 = False
|
|
||||||
else:
|
|
||||||
_FI_HAS_SM90_CUTLASS_MXFP4 = False
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -59,26 +30,21 @@ _GROUP_SIZE = 32
|
|||||||
|
|
||||||
|
|
||||||
class Mxfp4FlashinferCutlassMoEMethod:
|
class Mxfp4FlashinferCutlassMoEMethod:
|
||||||
"""DeepSeek-V4 W4A16 MXFP4 MoE via FlashInfer's SM90 mixed-input cutlass
|
"""FlashInfer MXFP4 MoE: W4A16 on SM90 and W4A8 on SM120."""
|
||||||
grouped GEMM. The fused kernel does GEMM1 + clamped SwiGLU + GEMM2 in one
|
|
||||||
call after a one-shot weight/scale interleave at load time."""
|
|
||||||
|
|
||||||
def __init__(self, fp8_method, prefix: str):
|
def __init__(self, fp8_method, prefix: str):
|
||||||
if not _FI_HAS_SM90_CUTLASS_MXFP4:
|
if not is_flashinfer_available():
|
||||||
raise RuntimeError(
|
raise RuntimeError("Mxfp4FlashinferCutlassMoEMethod requires FlashInfer.")
|
||||||
"Mxfp4FlashinferCutlassMoEMethod requires FlashInfer >= 0.6.11 "
|
self._use_mxfp8_act_scaling = is_sm120_supported()
|
||||||
"(PR #3084 SM90 mixed-input helpers). Older builds lack "
|
|
||||||
"interleave_moe_{weights,scales}_for_sm90_mixed_gemm; "
|
|
||||||
"either upgrade flashinfer-python or fall back to "
|
|
||||||
"--moe-runner-backend marlin."
|
|
||||||
)
|
|
||||||
self._fp8 = fp8_method
|
self._fp8 = fp8_method
|
||||||
self.prefix = prefix
|
self.prefix = prefix
|
||||||
self._swiglu_alpha_tensor: torch.Tensor | None = None
|
|
||||||
self._swiglu_beta_tensor: torch.Tensor | None = None
|
|
||||||
self._swiglu_limit_tensor: torch.Tensor | None = None
|
self._swiglu_limit_tensor: torch.Tensor | None = None
|
||||||
|
self._mxfp4_weight_global_scale_tensor: torch.Tensor | None = None
|
||||||
|
|
||||||
# --- Lifecycle ---------------------------------------------------------
|
@property
|
||||||
|
def load_up_proj_weight_first(self) -> bool:
|
||||||
|
"""Load W13 directly as ``[up; gate]`` for FlashInfer CUTLASS."""
|
||||||
|
return True
|
||||||
|
|
||||||
def create_weights(
|
def create_weights(
|
||||||
self,
|
self,
|
||||||
@@ -89,11 +55,7 @@ class Mxfp4FlashinferCutlassMoEMethod:
|
|||||||
params_dtype,
|
params_dtype,
|
||||||
**extra_weight_attrs,
|
**extra_weight_attrs,
|
||||||
):
|
):
|
||||||
# SM90 mixed-input GEMM: contraction dim K must be a multiple of 128
|
# Both CUTLASS paths require dimensions aligned to 128.
|
||||||
# (interleave factor = 128 / group_size = 4). For DSv4 (hidden=7168,
|
|
||||||
# inter=2048) both are already multiples of 128; we assert rather than
|
|
||||||
# silently pad here, since padding the FP8-base buffers in-place would
|
|
||||||
# require deeper changes.
|
|
||||||
if hidden_size % 128 != 0 or intermediate_size_per_partition % 128 != 0:
|
if hidden_size % 128 != 0 or intermediate_size_per_partition % 128 != 0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Mxfp4FlashinferCutlassMoEMethod requires hidden_size and "
|
"Mxfp4FlashinferCutlassMoEMethod requires hidden_size and "
|
||||||
@@ -101,14 +63,14 @@ class Mxfp4FlashinferCutlassMoEMethod:
|
|||||||
f"(got hidden={hidden_size}, "
|
f"(got hidden={hidden_size}, "
|
||||||
f"intermediate={intermediate_size_per_partition})."
|
f"intermediate={intermediate_size_per_partition})."
|
||||||
)
|
)
|
||||||
# Raw weight shapes match what the fp8 base method allocates for fp4
|
# Keep checkpoint scales in native E8M0 instead of staging them as FP32.
|
||||||
# experts (uint8 4-bit packed weights, fp32 E8M0 scales). Delegate.
|
|
||||||
self._fp8.create_weights(
|
self._fp8.create_weights(
|
||||||
layer,
|
layer,
|
||||||
num_experts,
|
num_experts,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
intermediate_size_per_partition,
|
intermediate_size_per_partition,
|
||||||
params_dtype,
|
params_dtype,
|
||||||
|
fp4_scale_dtype=torch.float8_e8m0fnu,
|
||||||
**extra_weight_attrs,
|
**extra_weight_attrs,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -118,29 +80,21 @@ class Mxfp4FlashinferCutlassMoEMethod:
|
|||||||
|
|
||||||
self.moe_runner_config = moe_runner_config
|
self.moe_runner_config = moe_runner_config
|
||||||
|
|
||||||
# DSv4 uses standard SwiGLU plus a config-driven activation clamp.
|
E = layer.num_local_experts
|
||||||
# We pass all three (alpha, beta, limit) as explicit per-expert tensors
|
device = layer.w13_weight.device
|
||||||
# rather than mixing tensors with None: the cutlass SwiGLU kernel
|
if self._use_mxfp8_act_scaling:
|
||||||
# branches on whether each is None, and partial-None inputs land in
|
# FlashInfer's MXFP4 ABI requires a neutral per-expert global scale.
|
||||||
# less-tested code paths. ``alpha=1.0``, ``beta=0.0`` reproduce plain
|
self._mxfp4_weight_global_scale_tensor = torch.ones(
|
||||||
# ``silu(gate) * up``; ``limit`` enforces the activation clamp the
|
E, dtype=torch.float32, device=device
|
||||||
# checkpoint was trained with.
|
)
|
||||||
|
|
||||||
|
# FlashInfer defaults alpha/beta to 1/0, so DSv4 only supplies its clamp.
|
||||||
swiglu_limit = getattr(moe_runner_config, "swiglu_limit", None)
|
swiglu_limit = getattr(moe_runner_config, "swiglu_limit", None)
|
||||||
if swiglu_limit is not None:
|
if swiglu_limit is not None:
|
||||||
E = layer.num_local_experts
|
|
||||||
device = layer.w13_weight.device
|
|
||||||
self._swiglu_alpha_tensor = torch.ones(
|
|
||||||
E, dtype=torch.float32, device=device
|
|
||||||
)
|
|
||||||
self._swiglu_beta_tensor = torch.zeros(
|
|
||||||
E, dtype=torch.float32, device=device
|
|
||||||
)
|
|
||||||
self._swiglu_limit_tensor = torch.full(
|
self._swiglu_limit_tensor = torch.full(
|
||||||
(E,), float(swiglu_limit), dtype=torch.float32, device=device
|
(E,), float(swiglu_limit), dtype=torch.float32, device=device
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self._swiglu_alpha_tensor = None
|
|
||||||
self._swiglu_beta_tensor = None
|
|
||||||
self._swiglu_limit_tensor = None
|
self._swiglu_limit_tensor = None
|
||||||
|
|
||||||
# Register the fused func at runner construction so the FusedOpPool
|
# Register the fused func at runner construction so the FusedOpPool
|
||||||
@@ -150,67 +104,71 @@ class Mxfp4FlashinferCutlassMoEMethod:
|
|||||||
self.runner = MoeRunner(MoeRunnerBackend.FLASHINFER_MXFP4, moe_runner_config)
|
self.runner = MoeRunner(MoeRunnerBackend.FLASHINFER_MXFP4, moe_runner_config)
|
||||||
|
|
||||||
def process_weights_after_loading(self, layer: Module) -> None:
|
def process_weights_after_loading(self, layer: Module) -> None:
|
||||||
from sglang.srt.layers.quantization.utils import reorder_w1w3_to_w3w1
|
# Preserve the base FP4 post-load handling.
|
||||||
|
|
||||||
# Run the fp8 base hook first (ROCm normalization, mxfp8 requant, ...).
|
|
||||||
self._fp8.process_weights_after_loading(layer)
|
self._fp8.process_weights_after_loading(layer)
|
||||||
|
|
||||||
if getattr(layer, "_mega_moe_weights_built", False):
|
if getattr(layer, "_mega_moe_weights_built", False):
|
||||||
return
|
return
|
||||||
|
|
||||||
# cutlass_fused_moe expects fc1 in [w3; w1] = [up; gate] order, just
|
arch = "SM120" if self._use_mxfp8_act_scaling else "SM90"
|
||||||
# like the trtllm-gen path. The HF / FP8 loader emits [w1; w3].
|
|
||||||
w13, w13_s = reorder_w1w3_to_w3w1(
|
|
||||||
layer.w13_weight.data, layer.w13_weight_scale_inv.data
|
|
||||||
)
|
|
||||||
layer.w13_weight = Parameter(w13, requires_grad=False)
|
|
||||||
layer.w13_weight_scale_inv = Parameter(w13_s, requires_grad=False)
|
|
||||||
|
|
||||||
log_info_on_rank0(
|
log_info_on_rank0(
|
||||||
logger,
|
logger,
|
||||||
f"Preparing DSv4 MXFP4 experts for FlashInfer SM90 cutlass "
|
f"Preparing DSv4 MXFP4 experts for FlashInfer {arch} CUTLASS "
|
||||||
f"(layer: {self.prefix})...",
|
f"(layer: {self.prefix})...",
|
||||||
)
|
)
|
||||||
|
|
||||||
# FP8 base stores scales as fp32 numerical values (= 2**e). The
|
# FlashInfer consumes the raw bytes of the checkpoint's E8M0 scales.
|
||||||
# FlashInfer SM90 helper reads raw E8M0 bytes (uint8 with the
|
for name in ("w13_weight_scale_inv", "w2_weight_scale_inv"):
|
||||||
# exponent + 127 bias). Cast through float8_e8m0fnu to extract the
|
scale = getattr(layer, name)
|
||||||
# raw byte without losing the exponent.
|
if scale.dtype != torch.float8_e8m0fnu:
|
||||||
w13_scale_u8 = (
|
raise TypeError(
|
||||||
layer.w13_weight_scale_inv.data.to(torch.float8_e8m0fnu)
|
f"{name} must remain native E8M0 for FlashInfer MXFP4, "
|
||||||
.view(torch.uint8)
|
f"got {scale.dtype}."
|
||||||
.contiguous()
|
)
|
||||||
)
|
w13_scale_u8 = layer.w13_weight_scale_inv.data.view(torch.uint8)
|
||||||
w2_scale_u8 = (
|
w2_scale_u8 = layer.w2_weight_scale_inv.data.view(torch.uint8)
|
||||||
layer.w2_weight_scale_inv.data.to(torch.float8_e8m0fnu)
|
|
||||||
.view(torch.uint8)
|
|
||||||
.contiguous()
|
|
||||||
)
|
|
||||||
|
|
||||||
# C++ byte interleave on packed 4-bit weights.
|
if self._use_mxfp8_act_scaling:
|
||||||
w13_il = interleave_moe_weights_for_sm90_mixed_gemm(
|
from flashinfer import block_scale_interleave
|
||||||
layer.w13_weight.data.view(torch.uint8).contiguous(), "fp4"
|
|
||||||
)
|
|
||||||
w2_il = interleave_moe_weights_for_sm90_mixed_gemm(
|
|
||||||
layer.w2_weight.data.view(torch.uint8).contiguous(), "fp4"
|
|
||||||
)
|
|
||||||
# Pure-PyTorch reshape+permute on E8M0 block scales.
|
|
||||||
w13_s_il = interleave_moe_scales_for_sm90_mixed_gemm(
|
|
||||||
w13_scale_u8, group_size=_GROUP_SIZE
|
|
||||||
)
|
|
||||||
w2_s_il = interleave_moe_scales_for_sm90_mixed_gemm(
|
|
||||||
w2_scale_u8, group_size=_GROUP_SIZE
|
|
||||||
)
|
|
||||||
|
|
||||||
layer.w13_weight = Parameter(w13_il, requires_grad=False)
|
if (
|
||||||
layer.w2_weight = Parameter(w2_il, requires_grad=False)
|
not layer.w13_weight.is_contiguous()
|
||||||
layer.w13_weight_scale_inv = Parameter(w13_s_il, requires_grad=False)
|
or not layer.w2_weight.is_contiguous()
|
||||||
layer.w2_weight_scale_inv = Parameter(w2_s_il, requires_grad=False)
|
):
|
||||||
|
raise ValueError("SM120 FlashInfer MXFP4 weights must be contiguous.")
|
||||||
|
for scale_u8 in (w13_scale_u8, w2_scale_u8):
|
||||||
|
scale_u8.copy_(block_scale_interleave(scale_u8).reshape_as(scale_u8))
|
||||||
|
else:
|
||||||
|
from flashinfer.fused_moe import (
|
||||||
|
interleave_moe_scales_for_sm90_mixed_gemm,
|
||||||
|
interleave_moe_weights_for_sm90_mixed_gemm,
|
||||||
|
)
|
||||||
|
|
||||||
layer._dsv4_mxfp4_backend = "flashinfer_cutlass_sm90"
|
w13_il = interleave_moe_weights_for_sm90_mixed_gemm(
|
||||||
torch.cuda.empty_cache()
|
layer.w13_weight.data.view(torch.uint8).contiguous(), "fp4"
|
||||||
|
)
|
||||||
|
w2_il = interleave_moe_weights_for_sm90_mixed_gemm(
|
||||||
|
layer.w2_weight.data.view(torch.uint8).contiguous(), "fp4"
|
||||||
|
)
|
||||||
|
w13_s_il = interleave_moe_scales_for_sm90_mixed_gemm(
|
||||||
|
w13_scale_u8, group_size=_GROUP_SIZE
|
||||||
|
)
|
||||||
|
w2_s_il = interleave_moe_scales_for_sm90_mixed_gemm(
|
||||||
|
w2_scale_u8, group_size=_GROUP_SIZE
|
||||||
|
)
|
||||||
|
layer.w13_weight = Parameter(w13_il, requires_grad=False)
|
||||||
|
layer.w2_weight = Parameter(w2_il, requires_grad=False)
|
||||||
|
layer.w13_weight_scale_inv = Parameter(w13_s_il, requires_grad=False)
|
||||||
|
layer.w2_weight_scale_inv = Parameter(w2_s_il, requires_grad=False)
|
||||||
|
|
||||||
# --- Forward -----------------------------------------------------------
|
layer._dsv4_mxfp4_backend = (
|
||||||
|
"flashinfer_cutlass_sm120"
|
||||||
|
if self._use_mxfp8_act_scaling
|
||||||
|
else "flashinfer_cutlass_sm90"
|
||||||
|
)
|
||||||
|
# SM90 creates full-size interleaved copies; release old layouts per layer.
|
||||||
|
if not self._use_mxfp8_act_scaling:
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
def apply(
|
def apply(
|
||||||
self,
|
self,
|
||||||
@@ -221,26 +179,21 @@ class Mxfp4FlashinferCutlassMoEMethod:
|
|||||||
FlashInferCutlassMxfp4MoeQuantInfo,
|
FlashInferCutlassMxfp4MoeQuantInfo,
|
||||||
)
|
)
|
||||||
|
|
||||||
# DSv4 always feeds StandardDispatchOutput; the fused func tolerates
|
|
||||||
# bypassed too but we keep the strict check here as a contract guard.
|
|
||||||
topk_output = dispatch_output.topk_output
|
|
||||||
if not TopKOutputChecker.format_is_standard(topk_output):
|
|
||||||
raise ValueError(f"Unsupported topk output format: {topk_output.format}")
|
|
||||||
|
|
||||||
quant_info = FlashInferCutlassMxfp4MoeQuantInfo(
|
quant_info = FlashInferCutlassMxfp4MoeQuantInfo(
|
||||||
w13_weight=layer.w13_weight,
|
w13_weight=layer.w13_weight,
|
||||||
w2_weight=layer.w2_weight,
|
w2_weight=layer.w2_weight,
|
||||||
w13_weight_scale=layer.w13_weight_scale_inv,
|
w13_weight_scale=layer.w13_weight_scale_inv,
|
||||||
w2_weight_scale=layer.w2_weight_scale_inv,
|
w2_weight_scale=layer.w2_weight_scale_inv,
|
||||||
w13_bias=None, # DSv4 has no MoE expert bias.
|
mxfp4_weight_global_scale=self._mxfp4_weight_global_scale_tensor,
|
||||||
|
w13_bias=None,
|
||||||
w2_bias=None,
|
w2_bias=None,
|
||||||
swiglu_alpha=self._swiglu_alpha_tensor, # ones: standard SiLU gate
|
swiglu_alpha=None,
|
||||||
swiglu_beta=self._swiglu_beta_tensor, # zeros: standard up
|
swiglu_beta=None,
|
||||||
swiglu_limit=self._swiglu_limit_tensor,
|
swiglu_limit=self._swiglu_limit_tensor,
|
||||||
moe_tp_size=layer.moe_tp_size,
|
moe_tp_size=layer.moe_tp_size,
|
||||||
moe_tp_rank=layer.moe_tp_rank,
|
moe_tp_rank=layer.moe_tp_rank,
|
||||||
moe_ep_size=layer.moe_ep_size,
|
moe_ep_size=layer.moe_ep_size,
|
||||||
moe_ep_rank=layer.moe_ep_rank,
|
moe_ep_rank=layer.moe_ep_rank,
|
||||||
padded_hidden=None, # DSv4 hidden_size is already a multiple of 128.
|
padded_hidden=None,
|
||||||
)
|
)
|
||||||
return self.runner.run(dispatch_output, quant_info)
|
return self.runner.run(dispatch_output, quant_info)
|
||||||
|
|||||||
@@ -71,21 +71,23 @@ class Mxfp4MarlinMoEMethod:
|
|||||||
layer.register_parameter("w2_weight", w2_weight)
|
layer.register_parameter("w2_weight", w2_weight)
|
||||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||||
|
|
||||||
|
# Store loader scales in E8M0; uint8 127 encodes 1.0.
|
||||||
|
def _e8m0_ones(*shape: int) -> torch.Tensor:
|
||||||
|
return torch.full(shape, 127, dtype=torch.uint8).view(torch.float8_e8m0fnu)
|
||||||
|
|
||||||
w13_weight_scale = torch.nn.Parameter(
|
w13_weight_scale = torch.nn.Parameter(
|
||||||
torch.ones(
|
_e8m0_ones(
|
||||||
num_experts,
|
num_experts,
|
||||||
2 * intermediate_size_per_partition,
|
2 * intermediate_size_per_partition,
|
||||||
hidden_size // fp4_block_k,
|
hidden_size // fp4_block_k,
|
||||||
dtype=torch.float32,
|
|
||||||
),
|
),
|
||||||
requires_grad=False,
|
requires_grad=False,
|
||||||
)
|
)
|
||||||
w2_weight_scale = torch.nn.Parameter(
|
w2_weight_scale = torch.nn.Parameter(
|
||||||
torch.ones(
|
_e8m0_ones(
|
||||||
num_experts,
|
num_experts,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
intermediate_size_per_partition // fp4_block_k,
|
intermediate_size_per_partition // fp4_block_k,
|
||||||
dtype=torch.float32,
|
|
||||||
),
|
),
|
||||||
requires_grad=False,
|
requires_grad=False,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -4470,6 +4470,8 @@ class ServerArgs:
|
|||||||
envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.set(False)
|
envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.set(False)
|
||||||
envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.set(False)
|
envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.set(False)
|
||||||
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.set(True)
|
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.set(True)
|
||||||
|
# Prefer TileLang over the Torch fallback.
|
||||||
|
envs.SGLANG_OPT_USE_TILELANG_INDEXER.set(True)
|
||||||
elif is_hip():
|
elif is_hip():
|
||||||
envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.set(False)
|
envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.set(False)
|
||||||
envs.SGLANG_OPT_USE_FUSED_COMPRESS.set(True)
|
envs.SGLANG_OPT_USE_FUSED_COMPRESS.set(True)
|
||||||
|
|||||||
@@ -0,0 +1,251 @@
|
|||||||
|
"""SM120 FlashInfer MXFP8-by-MXFP4 MoE integration test."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import builtins
|
||||||
|
import importlib
|
||||||
|
import sys
|
||||||
|
from contextlib import nullcontext
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=120, stage="base-b", runner_config="1-gpu-large")
|
||||||
|
|
||||||
|
|
||||||
|
def _random_weights(num_experts: int, hidden: int, intermediate: int):
|
||||||
|
generator = torch.Generator(device="cuda").manual_seed(0)
|
||||||
|
w13 = torch.randint(
|
||||||
|
-128,
|
||||||
|
128,
|
||||||
|
(num_experts, 2 * intermediate, hidden // 2),
|
||||||
|
dtype=torch.int8,
|
||||||
|
device="cuda",
|
||||||
|
generator=generator,
|
||||||
|
)
|
||||||
|
w2 = torch.randint(
|
||||||
|
-128,
|
||||||
|
128,
|
||||||
|
(num_experts, hidden, intermediate // 2),
|
||||||
|
dtype=torch.int8,
|
||||||
|
device="cuda",
|
||||||
|
generator=generator,
|
||||||
|
)
|
||||||
|
w13_scale_u8 = torch.randint(
|
||||||
|
125,
|
||||||
|
130,
|
||||||
|
(num_experts, 2 * intermediate, hidden // 32),
|
||||||
|
dtype=torch.uint8,
|
||||||
|
device="cuda",
|
||||||
|
generator=generator,
|
||||||
|
)
|
||||||
|
w2_scale_u8 = torch.randint(
|
||||||
|
125,
|
||||||
|
130,
|
||||||
|
(num_experts, hidden, intermediate // 32),
|
||||||
|
dtype=torch.uint8,
|
||||||
|
device="cuda",
|
||||||
|
generator=generator,
|
||||||
|
)
|
||||||
|
return (
|
||||||
|
w13,
|
||||||
|
w2,
|
||||||
|
w13_scale_u8.view(torch.float8_e8m0fnu),
|
||||||
|
w2_scale_u8.view(torch.float8_e8m0fnu),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_cutlass_adapter_import_does_not_require_flashinfer(monkeypatch):
|
||||||
|
module_name = "sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe"
|
||||||
|
# Load the package before blocking FlashInfer so this test isolates the
|
||||||
|
# adapter import exercised by non-CUDA backends.
|
||||||
|
importlib.import_module("sglang.srt.layers.quantization")
|
||||||
|
cached_module = sys.modules.pop(module_name, None)
|
||||||
|
real_import = builtins.__import__
|
||||||
|
|
||||||
|
def import_without_flashinfer(name, *args, **kwargs):
|
||||||
|
if name == "flashinfer" or name.startswith("flashinfer."):
|
||||||
|
raise ModuleNotFoundError("No module named 'flashinfer'")
|
||||||
|
return real_import(name, *args, **kwargs)
|
||||||
|
|
||||||
|
monkeypatch.setattr(builtins, "__import__", import_without_flashinfer)
|
||||||
|
try:
|
||||||
|
module = importlib.import_module(module_name)
|
||||||
|
assert hasattr(module, "Mxfp4FlashinferCutlassMoEMethod")
|
||||||
|
finally:
|
||||||
|
sys.modules.pop(module_name, None)
|
||||||
|
if cached_module is not None:
|
||||||
|
sys.modules[module_name] = cached_module
|
||||||
|
|
||||||
|
|
||||||
|
def test_dsv4_sm120_load_contract(monkeypatch):
|
||||||
|
import sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe as adapter_module
|
||||||
|
|
||||||
|
monkeypatch.setattr(adapter_module, "is_sm120_supported", lambda: True)
|
||||||
|
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
class _Fp8Method:
|
||||||
|
def create_weights(self, *args, **kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
|
||||||
|
method = adapter_module.Mxfp4FlashinferCutlassMoEMethod(_Fp8Method(), "test")
|
||||||
|
method.create_weights(
|
||||||
|
SimpleNamespace(),
|
||||||
|
num_experts=4,
|
||||||
|
hidden_size=256,
|
||||||
|
intermediate_size_per_partition=256,
|
||||||
|
params_dtype=torch.bfloat16,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert method.load_up_proj_weight_first
|
||||||
|
assert captured["fp4_scale_dtype"] == torch.float8_e8m0fnu
|
||||||
|
|
||||||
|
|
||||||
|
def test_dsv4_sm120_matches_direct_flashinfer(monkeypatch):
|
||||||
|
if not torch.cuda.is_available():
|
||||||
|
pytest.skip("CUDA required")
|
||||||
|
if torch.cuda.get_device_capability()[0] != 12:
|
||||||
|
pytest.skip("SM120 required")
|
||||||
|
pytest.importorskip("flashinfer.fused_moe")
|
||||||
|
|
||||||
|
from flashinfer import block_scale_interleave, mxfp8_quantize
|
||||||
|
from flashinfer.fused_moe import cutlass_fused_moe
|
||||||
|
from flashinfer.fused_moe.core import ActivationType
|
||||||
|
|
||||||
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass as runner_module
|
||||||
|
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.standard import StandardDispatchOutput
|
||||||
|
from sglang.srt.layers.moe.topk import StandardTopKOutput
|
||||||
|
from sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe import (
|
||||||
|
Mxfp4FlashinferCutlassMoEMethod,
|
||||||
|
)
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
runner_module, "use_symmetric_memory", lambda *args, **kwargs: nullcontext()
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(runner_module, "is_allocation_symmetric", lambda: False)
|
||||||
|
monkeypatch.setattr(runner_module, "get_tp_group", lambda: None)
|
||||||
|
|
||||||
|
num_experts, hidden, intermediate = 4, 256, 256
|
||||||
|
w13, w2, w13_scale, w2_scale = _random_weights(num_experts, hidden, intermediate)
|
||||||
|
w1, w3 = w13.chunk(2, dim=1)
|
||||||
|
w1_scale, w3_scale = w13_scale.chunk(2, dim=1)
|
||||||
|
# Simulate FusedMoE's ``load_up_proj_weight_first`` loader contract.
|
||||||
|
w31 = torch.cat((w3, w1), dim=1)
|
||||||
|
w31_scale = torch.cat(
|
||||||
|
(w3_scale.view(torch.uint8), w1_scale.view(torch.uint8)),
|
||||||
|
dim=1,
|
||||||
|
).view(torch.float8_e8m0fnu)
|
||||||
|
layer = SimpleNamespace(
|
||||||
|
w13_weight=torch.nn.Parameter(w31.clone(), requires_grad=False),
|
||||||
|
w2_weight=torch.nn.Parameter(w2.clone(), requires_grad=False),
|
||||||
|
w13_weight_scale_inv=torch.nn.Parameter(w31_scale.clone(), requires_grad=False),
|
||||||
|
w2_weight_scale_inv=torch.nn.Parameter(w2_scale.clone(), requires_grad=False),
|
||||||
|
num_local_experts=num_experts,
|
||||||
|
moe_tp_size=1,
|
||||||
|
moe_tp_rank=0,
|
||||||
|
moe_ep_size=1,
|
||||||
|
moe_ep_rank=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
method = Mxfp4FlashinferCutlassMoEMethod(
|
||||||
|
SimpleNamespace(process_weights_after_loading=lambda layer: None), "test"
|
||||||
|
)
|
||||||
|
config = MoeRunnerConfig(
|
||||||
|
num_experts=num_experts,
|
||||||
|
num_local_experts=num_experts,
|
||||||
|
hidden_size=hidden,
|
||||||
|
intermediate_size_per_partition=intermediate,
|
||||||
|
top_k=2,
|
||||||
|
activation="silu",
|
||||||
|
is_gated=True,
|
||||||
|
swiglu_limit=10,
|
||||||
|
)
|
||||||
|
method.create_moe_runner(layer, config)
|
||||||
|
|
||||||
|
w13_parameter = layer.w13_weight
|
||||||
|
w2_parameter = layer.w2_weight
|
||||||
|
w13_scale_parameter = layer.w13_weight_scale_inv
|
||||||
|
w2_scale_parameter = layer.w2_weight_scale_inv
|
||||||
|
method.process_weights_after_loading(layer)
|
||||||
|
|
||||||
|
expected_w13_scale = block_scale_interleave(w31_scale.view(torch.uint8)).reshape_as(
|
||||||
|
w31_scale
|
||||||
|
)
|
||||||
|
expected_w2_scale = block_scale_interleave(w2_scale.view(torch.uint8)).reshape_as(
|
||||||
|
w2_scale
|
||||||
|
)
|
||||||
|
assert layer.w13_weight is w13_parameter
|
||||||
|
assert layer.w2_weight is w2_parameter
|
||||||
|
assert layer.w13_weight_scale_inv is w13_scale_parameter
|
||||||
|
assert layer.w2_weight_scale_inv is w2_scale_parameter
|
||||||
|
assert torch.equal(layer.w13_weight_scale_inv.view(torch.uint8), expected_w13_scale)
|
||||||
|
assert torch.equal(layer.w2_weight_scale_inv.view(torch.uint8), expected_w2_scale)
|
||||||
|
|
||||||
|
generator = torch.Generator(device="cuda").manual_seed(1)
|
||||||
|
x = (
|
||||||
|
torch.randn(
|
||||||
|
8,
|
||||||
|
hidden,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device="cuda",
|
||||||
|
generator=generator,
|
||||||
|
)
|
||||||
|
* 0.1
|
||||||
|
)
|
||||||
|
logits = torch.randn(
|
||||||
|
8,
|
||||||
|
num_experts,
|
||||||
|
dtype=torch.float32,
|
||||||
|
device="cuda",
|
||||||
|
generator=generator,
|
||||||
|
)
|
||||||
|
topk_weights, topk_ids = torch.topk(torch.softmax(logits, dim=-1), 2, dim=-1)
|
||||||
|
topk_weights /= topk_weights.sum(dim=-1, keepdim=True)
|
||||||
|
topk = StandardTopKOutput(topk_weights, topk_ids.to(torch.int32), logits)
|
||||||
|
dispatch_output = StandardDispatchOutput(x, None, topk)
|
||||||
|
|
||||||
|
actual = method.apply(layer, dispatch_output).hidden_states
|
||||||
|
|
||||||
|
x_quant, x_scale = mxfp8_quantize(
|
||||||
|
x,
|
||||||
|
is_sf_swizzled_layout=True,
|
||||||
|
alignment=32,
|
||||||
|
)
|
||||||
|
global_scale = torch.ones(num_experts, dtype=torch.float32, device="cuda")
|
||||||
|
swiglu_limit = torch.full((num_experts,), 10.0, dtype=torch.float32, device="cuda")
|
||||||
|
expected = torch.empty_like(x)
|
||||||
|
cutlass_fused_moe(
|
||||||
|
input=x_quant,
|
||||||
|
token_selected_experts=topk_ids.to(torch.int32),
|
||||||
|
token_final_scales=topk_weights,
|
||||||
|
fc1_expert_weights=layer.w13_weight.view(torch.int64),
|
||||||
|
fc2_expert_weights=layer.w2_weight.view(torch.int64),
|
||||||
|
output_dtype=torch.bfloat16,
|
||||||
|
quant_scales=[
|
||||||
|
layer.w13_weight_scale_inv.view(torch.int32),
|
||||||
|
global_scale,
|
||||||
|
layer.w2_weight_scale_inv.view(torch.int32),
|
||||||
|
global_scale,
|
||||||
|
],
|
||||||
|
input_sf=x_scale,
|
||||||
|
# Compare the adapter's implicit defaults against the old explicit
|
||||||
|
# alpha=1/beta=0 representation.
|
||||||
|
swiglu_alpha=torch.ones(num_experts, dtype=torch.float32, device="cuda"),
|
||||||
|
swiglu_beta=torch.zeros(num_experts, dtype=torch.float32, device="cuda"),
|
||||||
|
swiglu_limit=swiglu_limit,
|
||||||
|
use_mxfp8_act_scaling=True,
|
||||||
|
activation_type=ActivationType.Swiglu,
|
||||||
|
tune_max_num_tokens=8,
|
||||||
|
output=expected,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert torch.equal(actual, expected)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__, "-v"]))
|
||||||
@@ -404,8 +404,7 @@ def test_apply_sm90_cutlass_matches_flashinfer_direct(
|
|||||||
|
|
||||||
|
|
||||||
def _make_random_dsv4_mxfp4(num_experts, hidden, inter, seed=0):
|
def _make_random_dsv4_mxfp4(num_experts, hidden, inter, seed=0):
|
||||||
"""Mirrors the fp8 base method's allocation for fp4 experts: int8-packed
|
"""Create native checkpoint-style packed MXFP4 weights and E8M0 scales."""
|
||||||
4-bit weights, fp32 scales (containing 2**e values, not raw E8M0 bytes)."""
|
|
||||||
g = torch.Generator(device="cuda").manual_seed(seed)
|
g = torch.Generator(device="cuda").manual_seed(seed)
|
||||||
# int8 storage (signed) -- matches Fp8MoEMethod.create_weights for fp4_experts.
|
# int8 storage (signed) -- matches Fp8MoEMethod.create_weights for fp4_experts.
|
||||||
w13 = torch.randint(
|
w13 = torch.randint(
|
||||||
@@ -424,8 +423,7 @@ def _make_random_dsv4_mxfp4(num_experts, hidden, inter, seed=0):
|
|||||||
device="cuda",
|
device="cuda",
|
||||||
generator=g,
|
generator=g,
|
||||||
)
|
)
|
||||||
# fp32 scales whose bit pattern after .to(float8_e8m0fnu).view(uint8) lands
|
# Native E8M0 scales with exponents around 0 (= 2**0).
|
||||||
# in a sane E8M0 band -- generate exponents around 0 (= 2**0).
|
|
||||||
raw_e = torch.randint(
|
raw_e = torch.randint(
|
||||||
125,
|
125,
|
||||||
130,
|
130,
|
||||||
@@ -442,8 +440,8 @@ def _make_random_dsv4_mxfp4(num_experts, hidden, inter, seed=0):
|
|||||||
device="cuda",
|
device="cuda",
|
||||||
generator=g,
|
generator=g,
|
||||||
)
|
)
|
||||||
w13_s = raw_e.view(torch.float8_e8m0fnu).to(torch.float32)
|
w13_s = raw_e.view(torch.float8_e8m0fnu)
|
||||||
w2_s = raw_e2.view(torch.float8_e8m0fnu).to(torch.float32)
|
w2_s = raw_e2.view(torch.float8_e8m0fnu)
|
||||||
return w13, w2, w13_s, w2_s
|
return w13, w2, w13_s, w2_s
|
||||||
|
|
||||||
|
|
||||||
@@ -460,12 +458,11 @@ def test_dsv4_apply_matches_flashinfer_direct(
|
|||||||
):
|
):
|
||||||
"""End-to-end: SGLang's DSv4 ``Mxfp4FlashinferCutlassMoEMethod.apply``
|
"""End-to-end: SGLang's DSv4 ``Mxfp4FlashinferCutlassMoEMethod.apply``
|
||||||
output must match a direct FlashInfer ``cutlass_fused_moe`` call with
|
output must match a direct FlashInfer ``cutlass_fused_moe`` call with
|
||||||
the equivalent reorder + scale-cast + interleave applied manually."""
|
the equivalent native E8M0 scale/weight interleave applied manually."""
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass as fi_cutlass_mod
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass as fi_cutlass_mod
|
||||||
import sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe as ds_mod
|
import sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe as ds_mod
|
||||||
from sglang.srt.layers.quantization.utils import reorder_w1w3_to_w3w1
|
|
||||||
|
|
||||||
# Bypass symmetric-memory / TP-group stack in the new fused-func module
|
# Bypass symmetric-memory / TP-group stack in the new fused-func module
|
||||||
# (where DSv4 ``apply`` now dispatches the kernel call through).
|
# (where DSv4 ``apply`` now dispatches the kernel call through).
|
||||||
@@ -476,29 +473,31 @@ def test_dsv4_apply_matches_flashinfer_direct(
|
|||||||
monkeypatch.setattr(fi_cutlass_mod, "get_tp_group", lambda: None)
|
monkeypatch.setattr(fi_cutlass_mod, "get_tp_group", lambda: None)
|
||||||
|
|
||||||
w13, w2, w13_s, w2_s = _make_random_dsv4_mxfp4(num_experts, hidden, inter)
|
w13, w2, w13_s, w2_s = _make_random_dsv4_mxfp4(num_experts, hidden, inter)
|
||||||
|
w1, w3 = w13.chunk(2, dim=1)
|
||||||
|
w1_s, w3_s = w13_s.chunk(2, dim=1)
|
||||||
|
# Simulate FusedMoE's ``load_up_proj_weight_first`` loader contract.
|
||||||
|
w31 = torch.cat((w3, w1), dim=1)
|
||||||
|
w31_s = torch.cat(
|
||||||
|
(w3_s.view(torch.uint8), w1_s.view(torch.uint8)),
|
||||||
|
dim=1,
|
||||||
|
).view(torch.float8_e8m0fnu)
|
||||||
x = torch.randn(tokens, hidden, dtype=torch.bfloat16, device="cuda") * 0.1
|
x = torch.randn(tokens, hidden, dtype=torch.bfloat16, device="cuda") * 0.1
|
||||||
topk_w, topk_i = _make_topk(tokens, num_experts, top_k)
|
topk_w, topk_i = _make_topk(tokens, num_experts, top_k)
|
||||||
|
|
||||||
# ---- SGLang DSv4 path ----
|
# ---- SGLang DSv4 path ----
|
||||||
method = ds_mod.Mxfp4FlashinferCutlassMoEMethod.__new__(
|
|
||||||
ds_mod.Mxfp4FlashinferCutlassMoEMethod
|
|
||||||
)
|
|
||||||
method._fp8 = SimpleNamespace(
|
|
||||||
process_weights_after_loading=lambda layer: None,
|
|
||||||
)
|
|
||||||
method.prefix = "test"
|
|
||||||
# plain SiLU * up — all three SwiGLU scalars None (no clamp configured).
|
# plain SiLU * up — all three SwiGLU scalars None (no clamp configured).
|
||||||
method._swiglu_alpha_tensor = None
|
method = ds_mod.Mxfp4FlashinferCutlassMoEMethod(
|
||||||
method._swiglu_beta_tensor = None
|
SimpleNamespace(process_weights_after_loading=lambda layer: None),
|
||||||
method._swiglu_limit_tensor = None
|
"test",
|
||||||
|
)
|
||||||
# Wire the unified MoeRunner -> flashinfer_mxfp4 fused func that
|
# Wire the unified MoeRunner -> flashinfer_mxfp4 fused func that
|
||||||
# ``apply`` now dispatches through.
|
# ``apply`` now dispatches through.
|
||||||
method.runner = _build_flashinfer_mxfp4_runner(num_experts, hidden, inter)
|
method.runner = _build_flashinfer_mxfp4_runner(num_experts, hidden, inter)
|
||||||
|
|
||||||
layer = _MockLayer()
|
layer = _MockLayer()
|
||||||
layer.w13_weight = torch.nn.Parameter(w13.clone(), requires_grad=False)
|
layer.w13_weight = torch.nn.Parameter(w31.clone(), requires_grad=False)
|
||||||
layer.w2_weight = torch.nn.Parameter(w2.clone(), requires_grad=False)
|
layer.w2_weight = torch.nn.Parameter(w2.clone(), requires_grad=False)
|
||||||
layer.w13_weight_scale_inv = torch.nn.Parameter(w13_s.clone(), requires_grad=False)
|
layer.w13_weight_scale_inv = torch.nn.Parameter(w31_s.clone(), requires_grad=False)
|
||||||
layer.w2_weight_scale_inv = torch.nn.Parameter(w2_s.clone(), requires_grad=False)
|
layer.w2_weight_scale_inv = torch.nn.Parameter(w2_s.clone(), requires_grad=False)
|
||||||
layer.num_local_experts = num_experts
|
layer.num_local_experts = num_experts
|
||||||
layer.moe_tp_size = 1
|
layer.moe_tp_size = 1
|
||||||
@@ -513,11 +512,10 @@ def test_dsv4_apply_matches_flashinfer_direct(
|
|||||||
).hidden_states
|
).hidden_states
|
||||||
|
|
||||||
# ---- Direct FlashInfer reference ----
|
# ---- Direct FlashInfer reference ----
|
||||||
w13_re, w13_s_re = reorder_w1w3_to_w3w1(w13, w13_s)
|
w13_s_u8 = w31_s.view(torch.uint8)
|
||||||
w13_s_u8 = w13_s_re.to(torch.float8_e8m0fnu).view(torch.uint8).contiguous()
|
w2_s_u8 = w2_s.view(torch.uint8)
|
||||||
w2_s_u8 = w2_s.to(torch.float8_e8m0fnu).view(torch.uint8).contiguous()
|
|
||||||
ref_w13 = interleave_moe_weights_for_sm90_mixed_gemm(
|
ref_w13 = interleave_moe_weights_for_sm90_mixed_gemm(
|
||||||
w13_re.view(torch.uint8).contiguous(), "fp4"
|
w31.view(torch.uint8).contiguous(), "fp4"
|
||||||
)
|
)
|
||||||
ref_w2 = interleave_moe_weights_for_sm90_mixed_gemm(
|
ref_w2 = interleave_moe_weights_for_sm90_mixed_gemm(
|
||||||
w2.view(torch.uint8).contiguous(), "fp4"
|
w2.view(torch.uint8).contiguous(), "fp4"
|
||||||
|
|||||||
@@ -521,6 +521,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
|||||||
"llama",
|
"llama",
|
||||||
dllm_algorithm="LowConfidence",
|
dllm_algorithm="LowConfidence",
|
||||||
disable_radix_cache=True,
|
disable_radix_cache=True,
|
||||||
|
attention_backend="triton",
|
||||||
)
|
)
|
||||||
self.assertEqual(sa.attention_backend, "flashinfer") # materialized
|
self.assertEqual(sa.attention_backend, "flashinfer") # materialized
|
||||||
self.assertIn(
|
self.assertIn(
|
||||||
@@ -821,7 +822,8 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
|||||||
|
|
||||||
with patch.object(overrides_module, "is_sm120_supported", return_value=True):
|
with patch.object(overrides_module, "is_sm120_supported", return_value=True):
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
_deepseek_v4_sm120_moe(_view()), {"moe_runner_backend": "marlin"}
|
_deepseek_v4_sm120_moe(_view()),
|
||||||
|
{"moe_runner_backend": "flashinfer_mxfp4"},
|
||||||
)
|
)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
_deepseek_v4_sm120_moe(_view(moe_runner_backend="triton")), {}
|
_deepseek_v4_sm120_moe(_view(moe_runner_backend="triton")), {}
|
||||||
|
|||||||
Reference in New Issue
Block a user