Enable MoE deferred finalize by default and drop its expert_weights dtype workaround (#33618)
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
c9506d023f
commit
beabc5949b
@@ -354,7 +354,7 @@ SGLang supports various environment variables that can be used to configure its
|
|||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_ENABLE_MOE_DEFERRED_FINALIZE</code></td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_ENABLE_MOE_DEFERRED_FINALIZE</code></td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Defer the MoE finalize step to overlap it with other work.</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Defer the MoE finalize step to overlap it with other work.</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>true</code></td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_PATCH_TOKENIZER</code></td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_PATCH_TOKENIZER</code></td>
|
||||||
|
|||||||
@@ -35,10 +35,10 @@
|
|||||||
* ``DeepseekV3MoE.forward``, and gives the downstream allreduce+rmsnorm
|
* ``DeepseekV3MoE.forward``, and gives the downstream allreduce+rmsnorm
|
||||||
* a clean PDL handoff.
|
* a clean PDL handoff.
|
||||||
*
|
*
|
||||||
* Expert-weight dtype is templated on ``TypeExpW`` so we support both the
|
* Expert-weight dtype is templated on ``TypeExpW`` so we accept both bf16
|
||||||
* bf16 and fp32 topk-weight paths (DSv3/K2.5 trtllm backends use fp32
|
* and fp32 topk weights. The trtllm deferred-finalize path always feeds bf16
|
||||||
* because their ``_routing_logits_dtype = torch.float32``; other backends
|
* (the trtllm-gen routing kernel emits bf16 for every routing method); fp32
|
||||||
* use bf16).
|
* is kept for callers that produce topk weights in fp32.
|
||||||
*
|
*
|
||||||
* Expert-weight scale convention: in our target backends
|
* Expert-weight scale convention: in our target backends
|
||||||
* (flashinfer trtllm nvfp4 + unquantized), ``apply_routed_scaling_factor_on_output``
|
* (flashinfer trtllm nvfp4 + unquantized), ``apply_routed_scaling_factor_on_output``
|
||||||
|
|||||||
@@ -1332,7 +1332,7 @@ class Envs:
|
|||||||
# Sglang Cache Dir
|
# Sglang Cache Dir
|
||||||
SGLANG_CACHE_DIR = EnvStr(os.path.expanduser("~/.cache/sglang"))
|
SGLANG_CACHE_DIR = EnvStr(os.path.expanduser("~/.cache/sglang"))
|
||||||
SGLANG_FLASHINFER_AUTOTUNE_CACHE = EnvBool(True)
|
SGLANG_FLASHINFER_AUTOTUNE_CACHE = EnvBool(True)
|
||||||
SGLANG_ENABLE_MOE_DEFERRED_FINALIZE = EnvBool(False)
|
SGLANG_ENABLE_MOE_DEFERRED_FINALIZE = EnvBool(True)
|
||||||
|
|
||||||
# Plugin system
|
# Plugin system
|
||||||
SGLANG_PLATFORM = EnvStr("")
|
SGLANG_PLATFORM = EnvStr("")
|
||||||
|
|||||||
@@ -1130,13 +1130,6 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
|
|||||||
result = trtllm_fp4_block_scale_moe(**moe_kwargs)
|
result = trtllm_fp4_block_scale_moe(**moe_kwargs)
|
||||||
if defer_finalize:
|
if defer_finalize:
|
||||||
gemm2_out, expert_weights, expanded_idx_to_permuted_idx = result[:3]
|
gemm2_out, expert_weights, expanded_idx_to_permuted_idx = result[:3]
|
||||||
# FIXME(kpham-sgl): flashinfer sizes this buffer from routing_logits
|
|
||||||
# dtype (fp32 in DSv3 decode) but always writes bf16 weights into it.
|
|
||||||
# Reinterpret the live bf16 prefix. Fix upstream alloc to drop this,
|
|
||||||
# tracking in https://github.com/flashinfer-ai/flashinfer/issues/3595
|
|
||||||
if expert_weights.dtype == torch.float32:
|
|
||||||
n, k = expert_weights.shape
|
|
||||||
expert_weights = expert_weights.view(torch.bfloat16).view(-1, k)[:n]
|
|
||||||
result = FlashInferTrtllmDeferredFinalizeOutput(
|
result = FlashInferTrtllmDeferredFinalizeOutput(
|
||||||
gemm2_out=gemm2_out,
|
gemm2_out=gemm2_out,
|
||||||
expert_weights=expert_weights,
|
expert_weights=expert_weights,
|
||||||
|
|||||||
Reference in New Issue
Block a user