Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
35 commits
Select commit Hold shift + click to select a range
b1ccde1
[Comm] Share the CuTe DSL AR fusion core and wire it for DeepSeek-V3 …
b8zhong Sep 4, 2026
16eb74d
[Comm] Trim the CuTe DSL fusion comments to the facts
b8zhong Sep 4, 2026
14296e3
[Comm] Drop the CuTe DSL workspace instance-count env var
b8zhong Sep 4, 2026
7a90133
[Comm] Delete dead surface in the CuTe DSL AR fusion
b8zhong Sep 4, 2026
6401718
[Comm] Bound the deferred MoE finalize by token count and absorb the …
b8zhong Sep 8, 2026
41af7ef
[Comm] Add the CuTe DSL AR fusion measurement record
b8zhong Sep 9, 2026
412b268
[Comm] Fix two wording defects in the measurement record
b8zhong Sep 9, 2026
e4e5d5e
[Comm] Report decode throughput instead of inter-token latency
b8zhong Sep 9, 2026
b3c396a
[Comm] Show the arms as a percentage gain over stock
b8zhong Sep 9, 2026
c25813f
Merge remote-tracking branch 'origin/main' into brayden/cutedsl-ar-fu…
Sep 12, 2026
148e75f
Drop two DSA/DCP tests that do not belong to this branch
Sep 12, 2026
4457499
Delete the unreachable should_use_finalize base stub
Sep 12, 2026
cca1513
Record one deferral flag instead of two
Sep 12, 2026
1a1536e
Delete Qwen3.5's tautological communicator check
Sep 12, 2026
29399f3
Separate incoming from outgoing post-MoE reduction eligibility
Sep 12, 2026
de189bb
Keep the deferred handoff inside the dual-stream op's tensor contract
Sep 12, 2026
7a342f9
Warn on the removed CuTe DSL fusion switch and fix the cookbook recipes
Sep 12, 2026
f1116b0
Drop the measurement record from the tree
Sep 12, 2026
bd63a99
Spell the fusion backend value cutedsl
Sep 12, 2026
8fa2317
Follow the repository rules in the CuTe DSL fusion
Sep 12, 2026
448f9bf
Move the dual-stream op contract test to the CPU fusion suite
Sep 12, 2026
4f1ad89
Merge remote-tracking branch 'origin/main' into brayden/cutedsl-ar-fu…
b8zhong Sep 16, 2026
92f9b83
Sync changes
Sep 17, 2026
bf920a7
Drop the cookbook changes
Sep 17, 2026
0e7b7b0
Restore the deferred MoE finalize token bound
Sep 16, 2026
2b365c8
Sync
Sep 17, 2026
0360803
Merge branch 'main' into brayden/cutedsl-ar-fusion-shared-core
mmangkad Sep 17, 2026
781d629
Merge remote-tracking branch 'origin/main' into brayden/cutedsl-ar-fu…
Sep 18, 2026
607a0c1
[MoE] Allow GlmMoeDsaForCausalLM for the FlashInfer MegaMoE backend
b8zhong Sep 18, 2026
21ba00e
Merge remote-tracking branch 'origin/main' into brayden/cutedsl-ar-fu…
Sep 19, 2026
61ab151
Merge remote-tracking branch 'origin/main' into brayden/cutedsl-ar-fu…
Sep 24, 2026
d0c1bc3
Address reviewers' feedback on the CuTe DSL AR fusion
Sep 24, 2026
e827f19
Address the second review round on the CuTe DSL AR fusion
Sep 24, 2026
1562d64
Run the GLM-5.2 NVFP4 TP MTP test on the CuTe DSL fusion
Sep 24, 2026
7edb3c7
Merge branch 'main' into brayden/cutedsl-ar-fusion-shared-core
mmangkad Sep 24, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion python/sglang/srt/arg_groups/fields/exec_.py
Original file line number Diff line number Diff line change
Expand Up @@ -630,7 +630,7 @@ class ExecComm(msgspec.Struct):
"Enforce disable FlashInfer allreduce fusion.",
] = False
flashinfer_allreduce_fusion_backend: A[
Optional[Literal["auto", "trtllm", "mnnvl"]],
Optional[Literal["auto", "trtllm", "mnnvl", "cutedsl"]],
Arg(
help=(
"Enable FlashInfer allreduce fusion and choose backend. "
Expand All @@ -641,6 +641,9 @@ class ExecComm(msgspec.Struct):
"'trtllm': available on single-node systems only. "
"'mnnvl': available on SM90 single-node systems and SM100/SM103 "
"single-node or multi-node systems via MNNVL fabric. "
"'cutedsl': Blackwell-only bf16 MNNVL CuTe DSL backend; also "
"fuses the MoE finalize and the shared-expert add into the "
"collective when the MoE runner can defer them. "
"Fuses allreduce with Residual + RMSNorm for supported MoE models."
),
resolvable=True,
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/arg_groups/moe_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -202,6 +202,7 @@ def validate_flashinfer_megamoe_model(server_args: Any) -> None:
"DeepseekV32ForCausalLM",
"DeepseekV4ForCausalLM",
"Glm4MoeForCausalLM",
"GlmMoeDsaForCausalLM",
"NemotronHForCausalLM",
"NemotronHPuzzleForCausalLM",
"Qwen2MoeForCausalLM",
Expand Down
14 changes: 0 additions & 14 deletions python/sglang/srt/arg_groups/overrides.py
Original file line number Diff line number Diff line change
Expand Up @@ -920,20 +920,6 @@ def _flashinfer_allreduce_fusion_auto_enable(view: Any) -> dict:
and view.nnodes == 1
and not view.disable_custom_all_reduce
)
if envs.SGLANG_FLASHINFER_MNNVL_CUTEDSL_AR_FUSION.get() and model_arch in {
"Qwen3_5MoeForCausalLM",
"Qwen3_5MoeForConditionalGeneration",
}:
# The Qwen backend owns one workspace for ordinary AR and MoE finalize;
# do not allocate or fall back to the legacy TRTLLM/MNNVL workspace.
if view.flashinfer_allreduce_fusion_backend is not None:
logger.warning(
"SGLANG_FLASHINFER_MNNVL_CUTEDSL_AR_FUSION owns both Qwen3.5 "
"AllReduce fusion patterns; suppressing the separately configured "
"--flashinfer-allreduce-fusion-backend=%s",
view.flashinfer_allreduce_fusion_backend,
)
return {"flashinfer_allreduce_fusion_backend": None}
if (
view.flashinfer_allreduce_fusion_backend is None
and model_arch in _FLASHINFER_ALLREDUCE_FUSION_ARCHS
Expand Down
4 changes: 4 additions & 0 deletions python/sglang/srt/distributed/bootstrap.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,8 +236,12 @@ def _set_all_reduce_flags() -> None:
set_custom_all_reduce(not get_exec().comm.disable_custom_all_reduce)
set_mscclpp_all_reduce(get_exec().comm.enable_mscclpp)
set_torch_symm_mem_all_reduce(get_exec().comm.enable_torch_symm_mem)
from sglang.srt.layers.flashinfer_comm_fusion import uses_cutedsl_ar_fusion

set_flashinfer_allreduce_only(
get_exec().comm.flashinfer_allreduce_fusion_backend is not None
# cutedsl has no legacy workspace for tagged groups to reduce over.
and not uses_cutedsl_ar_fusion()
)


Expand Down
22 changes: 15 additions & 7 deletions python/sglang/srt/environ.py
Original file line number Diff line number Diff line change
Expand Up @@ -1244,6 +1244,9 @@ class Envs:
SGLANG_NIXL_EP_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(128)
SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(128)
SGLANG_ENABLE_MOE_DEFERRED_FINALIZE = EnvBool(True)
# The [M*top_k, hidden] HBM round trip pays only at small M; 0 disables.
# GLM-5.2 (H=6144): neutral at 192 on B300 TP8, crossover 192-223 on GB300 TP4.
SGLANG_MOE_DEFERRED_FINALIZE_MAX_TOKENS = EnvInt(192)
# DeepSeek/GLM MoE (deepseek_v2.py): quantize the (dp-gathered) MoE input
# to per-token-group-128 fp8 ONCE and feed both the fused shared-expert
# GEMM (cutlass w8a8 linear) and the routed experts' triton fused runner,
Expand Down Expand Up @@ -1769,13 +1772,6 @@ class Envs:

# Qwen3.5 and GDN
SGLANG_ENABLE_GDN_DECODE_FUSED_PROJ_CONV = EnvBool(True)
SGLANG_TRACE_QWEN35_FINAL_NORM = EnvBool(False)
SGLANG_QWEN35_NATIVE_FINAL_NORM = EnvBool(False)
# One switch enables deferred MoE finalize and AR + residual + RMSNorm.
SGLANG_FLASHINFER_MNNVL_CUTEDSL_AR_FUSION = EnvBool(False)
# Distinct workspace configurations allowed in one process. Production
# uses one model/configuration per rank, so fail closed on accidental reuse.
SGLANG_FLASHINFER_MNNVL_CUTEDSL_AR_FUSION_MAX_INSTANCES = EnvInt(1)

# ===================================================================
# Plugin system
Expand Down Expand Up @@ -1887,10 +1883,22 @@ def apply(self, old_name: str):
# ad-hoc warnings. For a rename where the old name must keep working through a
# descriptor, use EnvBoolWithAlias / EnvIntWithAlias instead.
_DEPRECATED_ENVS: Dict[str, _DeprecatedEnv] = {
"SGLANG_FLASHINFER_MNNVL_CUTEDSL_AR_FUSION": _DeprecatedEnv(
note=(
"Pass --flashinfer-allreduce-fusion-backend cutedsl instead. "
"Without it an eligible model auto-enables the legacy mnnvl "
"backend rather than the CuTe DSL fusion."
)
),
"SGLANG_FLASHINFER_MNNVL_CUTEDSL_AR_FUSION_MAX_INSTANCES": _DeprecatedEnv(
note="One workspace per process is now an invariant, not a limit."
),
# Removed without replacement.
"SGLANG_ENABLE_CP_V2": _DeprecatedEnv(
note="Strategy-based prefill context parallelism is now the only generic implementation."
),
"SGLANG_TRACE_QWEN35_FINAL_NORM": _DeprecatedEnv(),
"SGLANG_QWEN35_NATIVE_FINAL_NORM": _DeprecatedEnv(),
"SGLANG_ENABLE_HICACHE_BUFFER_ANCHOR_LOCK": _DeprecatedEnv(
note="Buffer-mode anchor pinning is always on; set "
"SGLANG_HICACHE_BUFFER_ANCHOR_LOCK_CAP=0 to disable it."
Expand Down
35 changes: 32 additions & 3 deletions python/sglang/srt/layers/communicator.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,10 @@
is_enable_moe_cp_allgather,
moe_cp_all_gather_into_tensor,
)
from sglang.srt.layers.flashinfer_comm_fusion import is_flashinfer_allreduce_unavailable
from sglang.srt.layers.flashinfer_comm_fusion import (
is_flashinfer_allreduce_unavailable,
uses_cutedsl_ar_fusion,
)
from sglang.srt.layers.moe import (
can_merge_post_experts_all_reduce,
deferred_post_experts_all_reduce,
Expand Down Expand Up @@ -198,6 +201,8 @@ def apply_flashinfer_allreduce_fusion(batch_size: int):
and _is_flashinfer_available
and not is_dp_attention_enabled()
and get_exec().comm.flashinfer_allreduce_fusion_backend is not None
# cutedsl runs its own fused path from the fusion communicator.
and not uses_cutedsl_ar_fusion()
and not is_flashinfer_allreduce_unavailable()
# Symbolic size checks stay last: under Dynamo tracing they guard on
# the dynamic token dim, so statically-off configs must short-circuit
Expand Down Expand Up @@ -859,6 +864,14 @@ def prepare_attn(
post_residual_addition,
)

return self._finish_prepare_attn(
hidden_states=hidden_states,
residual=residual,
forward_batch=forward_batch,
)

def _finish_prepare_attn(self, hidden_states, residual, forward_batch):
"""Tail every prepare_attn path must run, or ``attn_inputs`` is unset."""
hidden_states = self._communicate_simple_fn(
hidden_states=hidden_states,
forward_batch=forward_batch,
Expand Down Expand Up @@ -953,6 +966,12 @@ def should_use_reduce_scatter(self, forward_batch: ForwardBatch):
return True
return False

def should_defer_moe_finalize(
self, forward_batch: ForwardBatch, m: int | None = None
) -> bool:
"""Whether the MoE may hand an unfinalized output to the next layer."""
return False

# NOTE: This function will cause torch recompilation
def should_fuse_mlp_allreduce_with_next_layer(
self, forward_batch: ForwardBatch
Expand Down Expand Up @@ -1046,11 +1065,13 @@ def scatter_mode_layouts(

class FfnExit:
"""One FFN's reduction decision. Inside the ``with`` block it is published as
``fuse_mlp_allreduce`` / ``mlp_reduce_scatter`` on ``get_forward()``."""
``fuse_mlp_allreduce`` / ``mlp_reduce_scatter`` / ``defer_moe_finalize`` on
``get_forward()``."""

__slots__ = (
"communicator",
"forward_batch",
"defer_moe_finalize",
"fuse_mlp_allreduce",
"mlp_reduce_scatter",
"_scope",
Expand All @@ -1059,13 +1080,17 @@ class FfnExit:
def __init__(self, communicator: LayerCommunicator, forward_batch: ForwardBatch):
self.communicator = communicator
self.forward_batch = forward_batch
self.defer_moe_finalize = communicator.should_defer_moe_finalize(forward_batch)
# Deferring implies fusing: a handoff skips the post-experts all-reduce.
self.fuse_mlp_allreduce = (
communicator.should_fuse_mlp_allreduce_with_next_layer(forward_batch)
self.defer_moe_finalize
or communicator.should_fuse_mlp_allreduce_with_next_layer(forward_batch)
)
self.mlp_reduce_scatter = communicator.should_use_reduce_scatter(forward_batch)
self._scope = get_forward().scoped(
fuse_mlp_allreduce=self.fuse_mlp_allreduce,
mlp_reduce_scatter=self.mlp_reduce_scatter,
defer_moe_finalize=self.defer_moe_finalize,
)

def __enter__(self) -> "FfnExit":
Expand All @@ -1079,6 +1104,10 @@ def finish(
self, hidden_states: torch.Tensor, residual: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Leave the reduction to the next layer's input norm, or postprocess."""
if not isinstance(hidden_states, torch.Tensor):
# A deferred MoE finalize handoff, consumed by the next prepare_attn.
assert self.defer_moe_finalize, "unrequested deferred MoE handoff"
return hidden_states, residual
if self.fuse_mlp_allreduce:
hidden_states._sglang_needs_allreduce_fusion = True
return hidden_states, residual
Expand Down
15 changes: 14 additions & 1 deletion python/sglang/srt/layers/flashinfer_comm_fusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,14 @@ def _resolve_backend(backend: str, is_multi_node: bool = False) -> str:
"FlashInfer allreduce fusion requires SM90 or SM10X NVIDIA GPUs."
)

if backend == "cutedsl":
if not get_platform().is_sm100:
raise ValueError(
"FlashInfer allreduce fusion cutedsl backend requires a "
"Blackwell system."
)
return backend

if backend == "auto":
if is_multi_node:
if get_platform().is_sm100:
Expand All @@ -72,6 +80,11 @@ def _resolve_backend(backend: str, is_multi_node: bool = False) -> str:
return backend


def uses_cutedsl_ar_fusion() -> bool:
"""Selected CuTe DSL owns both patterns, so the legacy workspace stands down."""
return get_exec().comm.flashinfer_allreduce_fusion_backend == "cutedsl"


def resolve_flashinfer_allreduce_fusion_backend() -> Optional[str]:
"""The fusion backend for this process, or None when fusion is off.

Expand Down Expand Up @@ -713,7 +726,7 @@ def ensure_workspace_initialized(
use_attn_tp_group: bool = True,
):
"""Ensure workspace is initialized."""
if _flashinfer_allreduce_unavailable:
if _flashinfer_allreduce_unavailable or uses_cutedsl_ar_fusion():
return False

if not is_flashinfer_available() or _flashinfer_comm is None:
Expand Down
Loading
Loading