From 05b91b941c6e66e0b71d9061ebb540e457daa86d Mon Sep 17 00:00:00 2001 From: Hao Phan Date: Wed, 30 Sep 2026 09:27:46 +0700 Subject: [PATCH 1/9] [Fix] Skip the deprecated getters' caller check while compiling In-package calls never mark a getter as warned, so every call from a torch.compile'd forward reaches sys._getframe, which Dynamo cannot trace. --- .../sglang/srt/distributed/parallel_state.py | 4 +++- test/registered/unit/test_runtime_context.py | 19 +++++++++++++++++++ 2 files changed, 22 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index d3277f43869a..d57c240f6454 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -3509,7 +3509,9 @@ def _warn_if_called_from_outside(name: str, replacement: str): def decorate(fn): @functools.wraps(fn) def wrapper(*args, **kwargs): - if name not in _ALREADY_WARNED: + # Dynamo cannot trace `sys._getframe`; + # in-package calls never mark `name` as warned, so they always reach it. + if not torch.compiler.is_compiling() and name not in _ALREADY_WARNED: caller = sys._getframe(1).f_globals.get("__name__", "") if not caller.startswith(_EXEMPT_CALLERS): _ALREADY_WARNED.add(name) diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index 48d7c6bbc4dc..e8c1278bbd70 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -2429,6 +2429,25 @@ def test_the_package_that_defines_them_is_not_warned_at(self): pass self.assertEqual([str(w.message) for w in seen], []) + def test_in_package_all_reduces_trace_under_fullgraph(self): + """The package's own all-reduces trace under fullgraph torch.compile.""" + import torch + + from sglang.srt.distributed import communication_op, parallel_state + + group = SimpleNamespace(all_reduce=lambda x: x * 2) + with ( + patch.object(parallel_state, "_ALREADY_WARNED", set()), + get_parallel().override(tp_group=group, attn_tp_group=group), + ): + for helper in ( + communication_op.tensor_model_parallel_all_reduce, + communication_op.attention_tensor_model_parallel_all_reduce, + ): + with self.subTest(helper.__name__): + compiled = torch.compile(helper, fullgraph=True, backend="eager") + self.assertEqual(compiled(torch.ones(2)).tolist(), [2.0, 2.0]) + def test_the_guard_would_notice_a_caller(self): self.assertTrue(self._callers("get_self_pp_group")) From 9d34d3c45aaa6d4a5fd81e2ac0c629b46ff1cc15 Mon Sep 17 00:00:00 2001 From: Hao Phan Date: Wed, 30 Sep 2026 13:51:38 +0700 Subject: [PATCH 2/9] [Fix] Build the per-forward layer boundary values as plain classes UnreducedOutput, DeclaredSum, Contribution, OwedOutput and ExitDecision are built (and Contribution mutated) on every forward inside the decoder stack. Dynamo cannot construct a msgspec.Struct, so a fullgraph torch.compile of that stack, the tc_piecewise prefill CUDA graph, failed. They become __slots__ classes with the same fields, defaults and constructor signatures, snapshot() rebuilds the unreduced output with its constructor, and the no-dataclasses rule records the exception. --- .claude/rules/no-dataclasses.md | 2 + .../sglang/srt/layers/layer_boundary/exit.py | 27 +++++++++--- .../srt/layers/layer_boundary/output.py | 22 +++++++--- .../layers/layer_boundary/residual/stream.py | 41 ++++++++++++++----- .../unit/layer_boundary/test_ffn_exit.py | 34 +++++++++++++++ .../layer_boundary/test_residual_stream.py | 27 ++++++++++++ 6 files changed, 130 insertions(+), 23 deletions(-) diff --git a/.claude/rules/no-dataclasses.md b/.claude/rules/no-dataclasses.md index 6a30511ec67c..9fa70b7435ae 100644 --- a/.claude/rules/no-dataclasses.md +++ b/.claude/rules/no-dataclasses.md @@ -21,3 +21,5 @@ class LoadSnapshot(msgspec.Struct): # prefer frozen= and omit_defaults=; kw_on `python/sglang/srt/managers/load_snapshot.py`. - New code only. Existing `@dataclass` is grandfathered — migrate opportunistically while editing the file, not in drive-by sweeps. +- Exception: values built or mutated inside a `torch.compile`d forward stay plain + `__slots__` classes; Dynamo cannot construct or mutate a `Struct`. diff --git a/python/sglang/srt/layers/layer_boundary/exit.py b/python/sglang/srt/layers/layer_boundary/exit.py index c4321f95db1a..12aa79f87868 100644 --- a/python/sglang/srt/layers/layer_boundary/exit.py +++ b/python/sglang/srt/layers/layer_boundary/exit.py @@ -18,7 +18,6 @@ from functools import partial from typing import Callable, Optional, Tuple -import msgspec import torch from sglang.srt.distributed import GroupCoordinator @@ -376,7 +375,9 @@ def _batch_allows_deferred_sum(forward_batch: ForwardBatch, boundary=None) -> bo return residual is not None and aiter_ar_fusion_applies(residual, forward_batch) -class ExitDecision(msgspec.Struct, frozen=True): +# A plain class, not msgspec.Struct; +# Dynamo cannot build a Struct inside a compiled layer. +class ExitDecision: """One decision shared by compute flags and the matching output completion. Fields: @@ -390,10 +391,24 @@ class ExitDecision(msgspec.Struct, frozen=True): Do not independently reselect completion after compute has used these flags. """ - defer_moe_finalize: bool - fuse_mlp_allreduce: bool - mlp_reduce_scatter: bool - complete: Callable[[torch.Tensor, torch.Tensor], Tuple] + __slots__ = ( + "defer_moe_finalize", + "fuse_mlp_allreduce", + "mlp_reduce_scatter", + "complete", + ) + + def __init__( + self, + defer_moe_finalize: bool, + fuse_mlp_allreduce: bool, + mlp_reduce_scatter: bool, + complete: Callable[[torch.Tensor, torch.Tensor], Tuple], + ): + self.defer_moe_finalize = defer_moe_finalize + self.fuse_mlp_allreduce = fuse_mlp_allreduce + self.mlp_reduce_scatter = mlp_reduce_scatter + self.complete = complete def _defer( diff --git a/python/sglang/srt/layers/layer_boundary/output.py b/python/sglang/srt/layers/layer_boundary/output.py index 935e17d75b67..be1a5057c272 100644 --- a/python/sglang/srt/layers/layer_boundary/output.py +++ b/python/sglang/srt/layers/layer_boundary/output.py @@ -37,7 +37,9 @@ class OutputTransform(msgspec.Struct, frozen=True): before_reduce_scatter: bool = False -class UnreducedOutput(msgspec.Struct, frozen=True): +# Per-forward values are plain classes, not msgspec.Struct; +# Dynamo cannot build a Struct inside a compiled layer. +class UnreducedOutput: """Internal adapter value describing an unfinished reduction. Fields: @@ -52,11 +54,19 @@ class UnreducedOutput(msgspec.Struct, frozen=True): required when reduce_to_dp_local is absent. """ - partial: torch.Tensor - group: Optional[GroupCoordinator] = None - # Under attention DP: the reduction that also brings ``partial`` back to this - # rank's tokens (a reduce-scatter, or an all-reduce then a scatter). - reduce_to_dp_local: Optional[Callable[[torch.Tensor], torch.Tensor]] = None + __slots__ = ("partial", "group", "reduce_to_dp_local") + + def __init__( + self, + partial: torch.Tensor, + group: Optional[GroupCoordinator] = None, + # Under attention DP: reduces ``partial`` and moves it to this rank's tokens, + # by a reduce-scatter or by an all-reduce then a scatter. + reduce_to_dp_local: Optional[Callable[[torch.Tensor], torch.Tensor]] = None, + ): + self.partial = partial + self.group = group + self.reduce_to_dp_local = reduce_to_dp_local def complete(self) -> torch.Tensor: """Complete the sum, on the destination rows when it moves them.""" diff --git a/python/sglang/srt/layers/layer_boundary/residual/stream.py b/python/sglang/srt/layers/layer_boundary/residual/stream.py index 586dc55a1eab..8c782be600f6 100644 --- a/python/sglang/srt/layers/layer_boundary/residual/stream.py +++ b/python/sglang/srt/layers/layer_boundary/residual/stream.py @@ -15,7 +15,6 @@ from typing import Optional, Union -import msgspec import torch from sglang.srt.layers.layer_boundary.layout import SumGroup, _sum_group @@ -23,16 +22,22 @@ from sglang.srt.layers.layer_boundary.residual import ResidualUpdate -class DeclaredSum(msgspec.Struct, frozen=True): +# Per-forward values are plain classes, not msgspec.Struct; +# Dynamo cannot build a Struct inside a compiled layer. +class DeclaredSum: """A sum every output of this producer owes to its declared input edge.""" - group: SumGroup + __slots__ = ("group",) + + def __init__(self, group: SumGroup): + self.group = group def complete(self, value): return _sum_group(self.group).all_reduce(value) -class Contribution(msgspec.Struct): +# Also mutated per forward, which Dynamo cannot do to a msgspec.Struct. +class Contribution: """Own a producer's output, residual update and outstanding completion. Fields: @@ -45,9 +50,17 @@ class Contribution(msgspec.Struct): Completing owed work clears owed but does not apply the residual update. """ - value: Optional[torch.Tensor] - update: ResidualUpdate - owed: Union[UnreducedOutput, DeclaredSum, DeferredFinalize, None] = None + __slots__ = ("value", "update", "owed") + + def __init__( + self, + value: Optional[torch.Tensor], + update: ResidualUpdate, + owed: Union[UnreducedOutput, DeclaredSum, DeferredFinalize, None] = None, + ): + self.value = value + self.update = update + self.owed = owed def for_boundary(self): # These forms are private inputs to the existing fused-kernel adapters. @@ -71,10 +84,13 @@ def release(self): self.owed = None -class OwedOutput(msgspec.Struct, frozen=True): +class OwedOutput: """Opaque model-facing handle. Only its boundary may read the contribution.""" - contribution: Contribution + __slots__ = ("contribution",) + + def __init__(self, contribution: Contribution): + self.contribution = contribution class ResidualStream: @@ -201,8 +217,11 @@ def snapshot(self, hidden): elif isinstance(pending.owed, DeclaredSum): value = pending.owed.complete(pending.value.clone()) elif isinstance(pending.owed, UnreducedOutput): - value = msgspec.structs.replace( - pending.owed, partial=pending.value.clone() + owed = pending.owed + value = UnreducedOutput( + pending.value.clone(), + group=owed.group, + reduce_to_dp_local=owed.reduce_to_dp_local, ).complete() else: raise NotImplementedError("a finalize handoff requires main-output capture") diff --git a/test/registered/unit/layer_boundary/test_ffn_exit.py b/test/registered/unit/layer_boundary/test_ffn_exit.py index da19337cfe0a..460fa55841d1 100644 --- a/test/registered/unit/layer_boundary/test_ffn_exit.py +++ b/test/registered/unit/layer_boundary/test_ffn_exit.py @@ -261,6 +261,40 @@ def test_deferral_implies_fusion_and_passes_the_handoff_through(self): ) self.assertIsInstance(hidden_states, UnreducedOutput) + def test_compiles_without_graph_breaks(self): + """The exit and the output it leaves trace under fullgraph torch.compile.""" + group = types.SimpleNamespace(all_reduce=lambda h: h * 3) + for fuse in (False, True): + with self.subTest(fuse=fuse): + communicator = stub_plan() + communicator.terminal = False + communicator.paths[BatchVariant.ORDINARY] = ordinary_steps( + OutputContract( + Layout(frozenset()), + group=SumGroup.TP, + may_defer_to_next=True, + ) + ) + output = communicator.output + output._defers_sum = lambda fb, steps, **_: fuse + output._skips_sum_for_reduce_scatter = lambda steps, dp: False + output.ffn_reduction_group = lambda steps: group + output._complete_now = lambda h, r, **_: (h + 1, r) + + def layer(hidden_states, residual): + stream = ResidualStream(residual) + with output.ffn_exit(None, stream=stream) as ffn_exit: + if get_forward().fuse_mlp_allreduce: + hidden_states = hidden_states * 2 + return stream.complete(ffn_exit.finish(hidden_states)) + + torch._dynamo.reset() + compiled = torch.compile(layer, backend="eager", fullgraph=True) + with patch_communicator("_batch_shards_over_cp", lambda fb: False): + hidden_states = compiled(self.hidden_states, self.residual) + expected = self.hidden_states * 6 if fuse else self.hidden_states + 1 + torch.testing.assert_close(hidden_states, expected) + class TestReduceOutput(CustomTestCase): def setUp(self): diff --git a/test/registered/unit/layer_boundary/test_residual_stream.py b/test/registered/unit/layer_boundary/test_residual_stream.py index f2941cefc5d4..9f711a9e4be6 100644 --- a/test/registered/unit/layer_boundary/test_residual_stream.py +++ b/test/registered/unit/layer_boundary/test_residual_stream.py @@ -423,6 +423,33 @@ def update_and_read(self, update, value, residual, norm, **kwargs): with self.assertRaises(RuntimeError): stream.input(hidden) + def test_layers_trace_under_fullgraph(self): + """A stream's record, complete, snapshot and write trace under fullgraph.""" + group = SimpleNamespace(all_reduce=lambda x: x * 2) + + def layers(hidden, residual): + stream = ResidualStream(residual) + for _ in range(2): + owed = stream.record(hidden, PLAIN_ADD, declared_sum=SumGroup.TP) + residual = stream.write(stream.complete(owed) + stream.residual) + owed = stream.record( + UnreducedOutput(residual * 3, group=group), PLAIN_ADD + ) + captured = stream.snapshot(owed) + _, residual = stream.input(owed) + hidden = stream.write(stream.complete(owed) + residual) + captured + return hidden + + with patch( + "sglang.srt.layers.layer_boundary.residual.stream._sum_group", + lambda sum_group: group, + ): + expected = layers(torch.ones(2, 4), torch.full((2, 4), 3.0)) + torch._dynamo.reset() + compiled = torch.compile(layers, fullgraph=True, backend="eager") + actual = compiled(torch.ones(2, 4), torch.full((2, 4), 3.0)) + torch.testing.assert_close(actual, expected) + class TestBatchStageOwnership(CustomTestCase): def test_terminal_norm_releases_layer_buffers(self): From 50f7dd19f6a8eea13502b577bf9b89dbf706fabd Mon Sep 17 00:00:00 2001 From: Hao Phan Date: Wed, 30 Sep 2026 11:38:23 +0700 Subject: [PATCH 3/9] [Fix] Build the deferred MoE finalize as a plain class MoeDeferredFinalize is built inside the MoE forward when the next layer's input takes over the finalize; like the other per-forward boundary values, DeferredFinalize and it become __slots__ classes Dynamo can construct. --- .../layers/layer_boundary/fusions/cutedsl.py | 34 +++++++++++++++---- .../srt/layers/layer_boundary/output.py | 4 ++- .../layer_boundary/test_residual_stream.py | 15 +++++++- 3 files changed, 44 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py b/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py index ed32bff45e95..ca5826c374d1 100644 --- a/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py +++ b/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py @@ -75,17 +75,37 @@ def _resolve_max_m(*, max_running_requests: int | None) -> int: return max(positive) -class MoeDeferredFinalize(DeferredFinalize, frozen=True): +# A plain class, not msgspec.Struct; +# Dynamo cannot build a Struct inside a compiled layer. +class MoeDeferredFinalize(DeferredFinalize): """Unfinalized routed output plus the separately gated shared contribution. The next layer's fused finalize + AR + add + norm takes it; ``finish`` is the MoE's own unfused tail, for any other reader.""" - routed_output: torch.Tensor - expert_weights: torch.Tensor - permuted_indices: torch.Tensor - gated_shared_output: torch.Tensor - m: int - finish: Callable[[], torch.Tensor] + __slots__ = ( + "routed_output", + "expert_weights", + "permuted_indices", + "gated_shared_output", + "m", + "finish", + ) + + def __init__( + self, + routed_output: torch.Tensor, + expert_weights: torch.Tensor, + permuted_indices: torch.Tensor, + gated_shared_output: torch.Tensor, + m: int, + finish: Callable[[], torch.Tensor], + ): + self.routed_output = routed_output + self.expert_weights = expert_weights + self.permuted_indices = permuted_indices + self.gated_shared_output = gated_shared_output + self.m = m + self.finish = finish def complete(self) -> torch.Tensor: return self.finish() diff --git a/python/sglang/srt/layers/layer_boundary/output.py b/python/sglang/srt/layers/layer_boundary/output.py index be1a5057c272..cc0493f7206f 100644 --- a/python/sglang/srt/layers/layer_boundary/output.py +++ b/python/sglang/srt/layers/layer_boundary/output.py @@ -75,13 +75,15 @@ def complete(self) -> torch.Tensor: return self.group.all_reduce(self.partial) -class DeferredFinalize(msgspec.Struct, frozen=True): +class DeferredFinalize: """A layer output that still owes work only its producer knows how to do (a MoE's finalize and sum), left for the next layer's input or for a terminal norm that accepts it (residual_batch.final_norm(finalize_norm=...)). A fused kernel there may do that work together with its own; anything else passes it through complete_owed(), which calls ``complete()``.""" + __slots__ = () + def complete(self) -> torch.Tensor: """Do the owed work, unfused, and return the complete output.""" raise NotImplementedError diff --git a/test/registered/unit/layer_boundary/test_residual_stream.py b/test/registered/unit/layer_boundary/test_residual_stream.py index 9f711a9e4be6..bc49c3003f4f 100644 --- a/test/registered/unit/layer_boundary/test_residual_stream.py +++ b/test/registered/unit/layer_boundary/test_residual_stream.py @@ -17,6 +17,7 @@ bind_entry, ) from sglang.srt.layers.layer_boundary.contracts import BatchVariant, StageKind +from sglang.srt.layers.layer_boundary.fusions.cutedsl import MoeDeferredFinalize from sglang.srt.layers.layer_boundary.output import UnreducedOutput from sglang.srt.layers.layer_boundary.residual.access import add_to_output from sglang.srt.layers.layer_boundary.residual.add_norm import PLAIN_ADD @@ -427,6 +428,16 @@ def test_layers_trace_under_fullgraph(self): """A stream's record, complete, snapshot and write trace under fullgraph.""" group = SimpleNamespace(all_reduce=lambda x: x * 2) + def finalize(rows): + return MoeDeferredFinalize( + routed_output=rows, + expert_weights=rows, + permuted_indices=rows, + gated_shared_output=rows, + m=rows.shape[0], + finish=lambda: rows * 5, + ) + def layers(hidden, residual): stream = ResidualStream(residual) for _ in range(2): @@ -437,7 +448,9 @@ def layers(hidden, residual): ) captured = stream.snapshot(owed) _, residual = stream.input(owed) - hidden = stream.write(stream.complete(owed) + residual) + captured + residual = stream.write(stream.complete(owed) + residual) + owed = stream.record(finalize(residual + captured), PLAIN_ADD) + hidden = stream.write(stream.complete(owed) + residual) return hidden with patch( From e1a748bcf0e6be382d12e6081b8ded9d029bbc43 Mon Sep 17 00:00:00 2001 From: Hao Phan Date: Mon, 5 Oct 2026 14:45:36 +0700 Subject: [PATCH 4/9] Drop the explanatory comments; the code says it --- python/sglang/srt/distributed/parallel_state.py | 2 -- python/sglang/srt/layers/layer_boundary/exit.py | 2 -- python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py | 2 -- python/sglang/srt/layers/layer_boundary/output.py | 6 ++---- python/sglang/srt/layers/layer_boundary/residual/stream.py | 3 --- 5 files changed, 2 insertions(+), 13 deletions(-) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index d57c240f6454..685bf5008ff3 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -3509,8 +3509,6 @@ def _warn_if_called_from_outside(name: str, replacement: str): def decorate(fn): @functools.wraps(fn) def wrapper(*args, **kwargs): - # Dynamo cannot trace `sys._getframe`; - # in-package calls never mark `name` as warned, so they always reach it. if not torch.compiler.is_compiling() and name not in _ALREADY_WARNED: caller = sys._getframe(1).f_globals.get("__name__", "") if not caller.startswith(_EXEMPT_CALLERS): diff --git a/python/sglang/srt/layers/layer_boundary/exit.py b/python/sglang/srt/layers/layer_boundary/exit.py index 12aa79f87868..d1729ec4bdcb 100644 --- a/python/sglang/srt/layers/layer_boundary/exit.py +++ b/python/sglang/srt/layers/layer_boundary/exit.py @@ -375,8 +375,6 @@ def _batch_allows_deferred_sum(forward_batch: ForwardBatch, boundary=None) -> bo return residual is not None and aiter_ar_fusion_applies(residual, forward_batch) -# A plain class, not msgspec.Struct; -# Dynamo cannot build a Struct inside a compiled layer. class ExitDecision: """One decision shared by compute flags and the matching output completion. diff --git a/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py b/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py index ca5826c374d1..d06c89a6aea3 100644 --- a/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py +++ b/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py @@ -75,8 +75,6 @@ def _resolve_max_m(*, max_running_requests: int | None) -> int: return max(positive) -# A plain class, not msgspec.Struct; -# Dynamo cannot build a Struct inside a compiled layer. class MoeDeferredFinalize(DeferredFinalize): """Unfinalized routed output plus the separately gated shared contribution. The next layer's fused finalize + AR + add + norm takes it; ``finish`` is diff --git a/python/sglang/srt/layers/layer_boundary/output.py b/python/sglang/srt/layers/layer_boundary/output.py index cc0493f7206f..713bd7073d09 100644 --- a/python/sglang/srt/layers/layer_boundary/output.py +++ b/python/sglang/srt/layers/layer_boundary/output.py @@ -37,8 +37,6 @@ class OutputTransform(msgspec.Struct, frozen=True): before_reduce_scatter: bool = False -# Per-forward values are plain classes, not msgspec.Struct; -# Dynamo cannot build a Struct inside a compiled layer. class UnreducedOutput: """Internal adapter value describing an unfinished reduction. @@ -60,8 +58,8 @@ def __init__( self, partial: torch.Tensor, group: Optional[GroupCoordinator] = None, - # Under attention DP: reduces ``partial`` and moves it to this rank's tokens, - # by a reduce-scatter or by an all-reduce then a scatter. + # Under attention DP: the reduction that also brings ``partial`` back to this + # rank's tokens (a reduce-scatter, or an all-reduce then a scatter). reduce_to_dp_local: Optional[Callable[[torch.Tensor], torch.Tensor]] = None, ): self.partial = partial diff --git a/python/sglang/srt/layers/layer_boundary/residual/stream.py b/python/sglang/srt/layers/layer_boundary/residual/stream.py index 8c782be600f6..cf1b20e73bad 100644 --- a/python/sglang/srt/layers/layer_boundary/residual/stream.py +++ b/python/sglang/srt/layers/layer_boundary/residual/stream.py @@ -22,8 +22,6 @@ from sglang.srt.layers.layer_boundary.residual import ResidualUpdate -# Per-forward values are plain classes, not msgspec.Struct; -# Dynamo cannot build a Struct inside a compiled layer. class DeclaredSum: """A sum every output of this producer owes to its declared input edge.""" @@ -36,7 +34,6 @@ def complete(self, value): return _sum_group(self.group).all_reduce(value) -# Also mutated per forward, which Dynamo cannot do to a msgspec.Struct. class Contribution: """Own a producer's output, residual update and outstanding completion. From c8e203e949024605ded99cf41c6859a34a04a9f9 Mon Sep 17 00:00:00 2001 From: Hao Phan Date: Mon, 5 Oct 2026 15:36:31 +0700 Subject: [PATCH 5/9] Drop the no-dataclasses rule exception --- .claude/rules/no-dataclasses.md | 2 -- 1 file changed, 2 deletions(-) diff --git a/.claude/rules/no-dataclasses.md b/.claude/rules/no-dataclasses.md index 9fa70b7435ae..6a30511ec67c 100644 --- a/.claude/rules/no-dataclasses.md +++ b/.claude/rules/no-dataclasses.md @@ -21,5 +21,3 @@ class LoadSnapshot(msgspec.Struct): # prefer frozen= and omit_defaults=; kw_on `python/sglang/srt/managers/load_snapshot.py`. - New code only. Existing `@dataclass` is grandfathered — migrate opportunistically while editing the file, not in drive-by sweeps. -- Exception: values built or mutated inside a `torch.compile`d forward stay plain - `__slots__` classes; Dynamo cannot construct or mutate a `Struct`. From 6f999c32b36fceda2cceca76619e54061626d991 Mon Sep 17 00:00:00 2001 From: Hao Phan Date: Mon, 5 Oct 2026 17:06:36 +0700 Subject: [PATCH 6/9] Run sparse attention as an eager break under breakable The breakable prefill graph called the sparse backend inline, so MiniMax-M3's sparse prefill attention was captured with the capture batch's metadata and replayed against real batches. Wrap the sparse op with eager_on_graph like the dense ops so it reruns against each replay batch. --- python/sglang/srt/layers/radix_attention.py | 14 +++-- .../unit/layers/test_radix_attention.py | 59 +++++++++++++++++++ 2 files changed, 68 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/layers/radix_attention.py b/python/sglang/srt/layers/radix_attention.py index 6840c21e2f08..2b81780a91bd 100644 --- a/python/sglang/srt/layers/radix_attention.py +++ b/python/sglang/srt/layers/radix_attention.py @@ -188,10 +188,6 @@ def forward( and (torch.compiler.is_compiling() or not _force_eager_attn.get()) ): if kwargs.get("idx_q") is not None: - if is_in_breakable_cuda_graph(): - return get_attn_backend().forward( - q, k, v, self, forward_batch, save_kv_cache, **kwargs - ) idx_q = kwargs["idx_q"] idx_k = kwargs["idx_k"] idx_v = kwargs.get("idx_v") @@ -199,7 +195,12 @@ def forward( (q.shape[0], self.tp_q_head_num * self.v_head_dim) ) idx_out = q.new_empty((q.shape[0], idx_q.shape[1] * idx_q.shape[2])) - unified_sparse_attention_with_output( + op = ( + breakable_unified_sparse_attention_with_output + if is_in_breakable_cuda_graph() + else unified_sparse_attention_with_output + ) + op( q, k, v, @@ -591,6 +592,9 @@ def unified_sparse_attention_with_output( breakable_unified_attention_with_output_and_lse = eager_on_graph(True)( unified_attention_with_output_and_lse ) +breakable_unified_sparse_attention_with_output = eager_on_graph(True)( + unified_sparse_attention_with_output +) def attention_with_output_extra_kwargs( diff --git a/test/registered/unit/layers/test_radix_attention.py b/test/registered/unit/layers/test_radix_attention.py index e2d7dc5a7459..82065d7ad7b4 100644 --- a/test/registered/unit/layers/test_radix_attention.py +++ b/test/registered/unit/layers/test_radix_attention.py @@ -181,6 +181,65 @@ def output_and_lse(*args, **kwargs): self.assertEqual(output.shape, query.shape) self.assertTrue(torch.all(output == 5)) + def test_sparse_attention_breaks_the_graph_under_breakable(self): + """Breakable replays captured segments with the capture batch's metadata + baked in, so sparse attention must run as an eager break. Called inline, + MiniMax-M3 replayed stale metadata and hit an illegal memory access.""" + layer = self._new_layer() + query = torch.zeros((4, 2, 3)) + idx_q = torch.zeros((4, 1, 3)) + op_names = { + False: "unified_sparse_attention_with_output", + True: "breakable_unified_sparse_attention_with_output", + } + + def fill_outputs(*args, **kwargs): + args[3].fill_(5) + args[4].fill_(7) + + for breakable in (False, True): + with self.subTest(breakable=breakable): + forward_batch = SimpleNamespace(forward_mode=ForwardMode.EXTEND) + with ExitStack() as stack: + stack.enter_context( + patch.object( + radix_attention_module, + "get_tc_piecewise_forward_context", + return_value=SimpleNamespace(), + ) + ) + stack.enter_context( + patch.object( + radix_attention_module, + "is_in_breakable_cuda_graph", + return_value=breakable, + ) + ) + backend = stack.enter_context( + patch.object(radix_attention_module, "get_attn_backend") + ) + backend.return_value.forward.return_value = ( + torch.zeros((4, 3)), + torch.zeros((4, 6)), + ) + mocks = { + name: stack.enter_context( + patch.object( + radix_attention_module, name, side_effect=fill_outputs + ) + ) + for name in op_names.values() + } + idx_out, attn_out = layer( + query, query, query, forward_batch, idx_q=idx_q, idx_k=query + ) + + backend.assert_not_called() + for name, mock in mocks.items(): + self.assertEqual(mock.call_count, int(name == op_names[breakable])) + self.assertTrue(torch.all(attn_out == 5)) + self.assertTrue(torch.all(idx_out == 7)) + def test_prefill_wrapper_opt_out_preserves_expanded_rows_and_batch(self): """Expanded attention rows must not be sliced using the runner's token count.""" layer = RadixAttention( From c50fd693ffca33ea7ef3a29189cd0a9dbff5189d Mon Sep 17 00:00:00 2001 From: Hao Phan Date: Mon, 5 Oct 2026 17:06:36 +0700 Subject: [PATCH 7/9] Default MiniMax-M3 prefill to the breakable CUDA graph --- python/sglang/srt/configs/model_config.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index b7d31afaef77..6b80a5631b97 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -2316,6 +2316,8 @@ def is_generation_model(model_architectures: List[str], is_embedding: bool = Fal "MuseGlimmerForConditionalGeneration", "KimiK3ForConditionalGeneration", "KimiK25ForConditionalGeneration", + "MiniMaxM3SparseForCausalLM", + "MiniMaxM3SparseForConditionalGeneration", ] if external_mm_model_arch := envs.SGLANG_EXTERNAL_MM_MODEL_ARCH.get(): From 5a1b3ba319a1853d2e7826195099628b2082a730 Mon Sep 17 00:00:00 2001 From: Hao Phan Date: Mon, 5 Oct 2026 22:39:27 +0700 Subject: [PATCH 8/9] Port the fullgraph FFN-exit test to main's exit flags FfnExit now publishes only defer_moe_finalize, and _sum_in_reduce_scatter replaced _skips_sum_for_reduce_scatter. --- test/registered/unit/layer_boundary/test_ffn_exit.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/test/registered/unit/layer_boundary/test_ffn_exit.py b/test/registered/unit/layer_boundary/test_ffn_exit.py index c884e322a0db..f39fc83211c2 100644 --- a/test/registered/unit/layer_boundary/test_ffn_exit.py +++ b/test/registered/unit/layer_boundary/test_ffn_exit.py @@ -269,8 +269,8 @@ def test_deferral_implies_fusion_and_passes_the_handoff_through(self): def test_compiles_without_graph_breaks(self): """The exit and the output it leaves trace under fullgraph torch.compile.""" group = types.SimpleNamespace(all_reduce=lambda h: h * 3) - for fuse in (False, True): - with self.subTest(fuse=fuse): + for defers in (False, True): + with self.subTest(defers=defers): communicator = stub_plan() communicator.terminal = False communicator.paths[BatchVariant.ORDINARY] = ordinary_steps( @@ -281,15 +281,15 @@ def test_compiles_without_graph_breaks(self): ) ) output = communicator.output - output._defers_sum = lambda fb, steps, **_: fuse - output._skips_sum_for_reduce_scatter = lambda steps, dp: False + output._defers_sum = lambda fb, steps, **_: defers + output._sum_in_reduce_scatter = lambda steps, dp_step: False output.ffn_reduction_group = lambda steps: group output._complete_now = lambda h, r, **_: (h + 1, r) def layer(hidden_states, residual): stream = ResidualStream(residual) with output.ffn_exit(None, stream=stream) as ffn_exit: - if get_forward().fuse_mlp_allreduce: + if get_forward().defer_moe_finalize: hidden_states = hidden_states * 2 return stream.complete(ffn_exit.finish(hidden_states)) @@ -297,7 +297,7 @@ def layer(hidden_states, residual): compiled = torch.compile(layer, backend="eager", fullgraph=True) with patch_communicator("_batch_shards_over_cp", lambda fb: False): hidden_states = compiled(self.hidden_states, self.residual) - expected = self.hidden_states * 6 if fuse else self.hidden_states + 1 + expected = self.hidden_states * 3 if defers else self.hidden_states + 1 torch.testing.assert_close(hidden_states, expected) From 0579e2052ab220f6cc24bfc05e4347f82866b6d1 Mon Sep 17 00:00:00 2001 From: Hao Phan Date: Mon, 5 Oct 2026 22:48:52 +0700 Subject: [PATCH 9/9] Drop the tc_piecewise fixes tc_piecewise is being removed (#41634) and M3 now defaults to breakable, so restore these files to main. --- .../sglang/srt/distributed/parallel_state.py | 2 +- .../sglang/srt/layers/layer_boundary/exit.py | 17 +++----- .../layers/layer_boundary/fusions/cutedsl.py | 32 ++++----------- .../srt/layers/layer_boundary/output.py | 24 ++++------- .../layers/layer_boundary/residual/stream.py | 38 +++++------------- .../unit/layer_boundary/test_ffn_exit.py | 34 ---------------- .../layer_boundary/test_residual_stream.py | 40 ------------------- test/registered/unit/test_runtime_context.py | 19 --------- 8 files changed, 31 insertions(+), 175 deletions(-) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 878b537ee964..f5441452aef8 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -3509,7 +3509,7 @@ def _warn_if_called_from_outside(name: str, replacement: str): def decorate(fn): @functools.wraps(fn) def wrapper(*args, **kwargs): - if not torch.compiler.is_compiling() and name not in _ALREADY_WARNED: + if name not in _ALREADY_WARNED: caller = sys._getframe(1).f_globals.get("__name__", "") if not caller.startswith(_EXEMPT_CALLERS): _ALREADY_WARNED.add(name) diff --git a/python/sglang/srt/layers/layer_boundary/exit.py b/python/sglang/srt/layers/layer_boundary/exit.py index 3f6bb95f92be..35585fcd77d6 100644 --- a/python/sglang/srt/layers/layer_boundary/exit.py +++ b/python/sglang/srt/layers/layer_boundary/exit.py @@ -18,6 +18,7 @@ from functools import partial from typing import Callable, Optional, Tuple +import msgspec import torch from sglang.srt.distributed import GroupCoordinator @@ -387,7 +388,7 @@ def _batch_allows_deferred_sum(forward_batch: ForwardBatch, boundary=None) -> bo return residual is not None and aiter_ar_fusion_applies(residual, forward_batch) -class ExitDecision: +class ExitDecision(msgspec.Struct, frozen=True): """One decision for an FFN output, made before compute runs. Fields: @@ -399,17 +400,9 @@ class ExitDecision: Do not independently reselect completion after compute has run. """ - __slots__ = ("defer_moe_finalize", "sum_in_reduce_scatter", "complete") - - def __init__( - self, - defer_moe_finalize: bool, - sum_in_reduce_scatter: bool, - complete: Callable[[torch.Tensor, torch.Tensor], Tuple], - ): - self.defer_moe_finalize = defer_moe_finalize - self.sum_in_reduce_scatter = sum_in_reduce_scatter - self.complete = complete + defer_moe_finalize: bool + sum_in_reduce_scatter: bool + complete: Callable[[torch.Tensor, torch.Tensor], Tuple] def _defer( diff --git a/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py b/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py index d06c89a6aea3..ed32bff45e95 100644 --- a/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py +++ b/python/sglang/srt/layers/layer_boundary/fusions/cutedsl.py @@ -75,35 +75,17 @@ def _resolve_max_m(*, max_running_requests: int | None) -> int: return max(positive) -class MoeDeferredFinalize(DeferredFinalize): +class MoeDeferredFinalize(DeferredFinalize, frozen=True): """Unfinalized routed output plus the separately gated shared contribution. The next layer's fused finalize + AR + add + norm takes it; ``finish`` is the MoE's own unfused tail, for any other reader.""" - __slots__ = ( - "routed_output", - "expert_weights", - "permuted_indices", - "gated_shared_output", - "m", - "finish", - ) - - def __init__( - self, - routed_output: torch.Tensor, - expert_weights: torch.Tensor, - permuted_indices: torch.Tensor, - gated_shared_output: torch.Tensor, - m: int, - finish: Callable[[], torch.Tensor], - ): - self.routed_output = routed_output - self.expert_weights = expert_weights - self.permuted_indices = permuted_indices - self.gated_shared_output = gated_shared_output - self.m = m - self.finish = finish + routed_output: torch.Tensor + expert_weights: torch.Tensor + permuted_indices: torch.Tensor + gated_shared_output: torch.Tensor + m: int + finish: Callable[[], torch.Tensor] def complete(self) -> torch.Tensor: return self.finish() diff --git a/python/sglang/srt/layers/layer_boundary/output.py b/python/sglang/srt/layers/layer_boundary/output.py index 713bd7073d09..935e17d75b67 100644 --- a/python/sglang/srt/layers/layer_boundary/output.py +++ b/python/sglang/srt/layers/layer_boundary/output.py @@ -37,7 +37,7 @@ class OutputTransform(msgspec.Struct, frozen=True): before_reduce_scatter: bool = False -class UnreducedOutput: +class UnreducedOutput(msgspec.Struct, frozen=True): """Internal adapter value describing an unfinished reduction. Fields: @@ -52,19 +52,11 @@ class UnreducedOutput: required when reduce_to_dp_local is absent. """ - __slots__ = ("partial", "group", "reduce_to_dp_local") - - def __init__( - self, - partial: torch.Tensor, - group: Optional[GroupCoordinator] = None, - # Under attention DP: the reduction that also brings ``partial`` back to this - # rank's tokens (a reduce-scatter, or an all-reduce then a scatter). - reduce_to_dp_local: Optional[Callable[[torch.Tensor], torch.Tensor]] = None, - ): - self.partial = partial - self.group = group - self.reduce_to_dp_local = reduce_to_dp_local + partial: torch.Tensor + group: Optional[GroupCoordinator] = None + # Under attention DP: the reduction that also brings ``partial`` back to this + # rank's tokens (a reduce-scatter, or an all-reduce then a scatter). + reduce_to_dp_local: Optional[Callable[[torch.Tensor], torch.Tensor]] = None def complete(self) -> torch.Tensor: """Complete the sum, on the destination rows when it moves them.""" @@ -73,15 +65,13 @@ def complete(self) -> torch.Tensor: return self.group.all_reduce(self.partial) -class DeferredFinalize: +class DeferredFinalize(msgspec.Struct, frozen=True): """A layer output that still owes work only its producer knows how to do (a MoE's finalize and sum), left for the next layer's input or for a terminal norm that accepts it (residual_batch.final_norm(finalize_norm=...)). A fused kernel there may do that work together with its own; anything else passes it through complete_owed(), which calls ``complete()``.""" - __slots__ = () - def complete(self) -> torch.Tensor: """Do the owed work, unfused, and return the complete output.""" raise NotImplementedError diff --git a/python/sglang/srt/layers/layer_boundary/residual/stream.py b/python/sglang/srt/layers/layer_boundary/residual/stream.py index cf1b20e73bad..586dc55a1eab 100644 --- a/python/sglang/srt/layers/layer_boundary/residual/stream.py +++ b/python/sglang/srt/layers/layer_boundary/residual/stream.py @@ -15,6 +15,7 @@ from typing import Optional, Union +import msgspec import torch from sglang.srt.layers.layer_boundary.layout import SumGroup, _sum_group @@ -22,19 +23,16 @@ from sglang.srt.layers.layer_boundary.residual import ResidualUpdate -class DeclaredSum: +class DeclaredSum(msgspec.Struct, frozen=True): """A sum every output of this producer owes to its declared input edge.""" - __slots__ = ("group",) - - def __init__(self, group: SumGroup): - self.group = group + group: SumGroup def complete(self, value): return _sum_group(self.group).all_reduce(value) -class Contribution: +class Contribution(msgspec.Struct): """Own a producer's output, residual update and outstanding completion. Fields: @@ -47,17 +45,9 @@ class Contribution: Completing owed work clears owed but does not apply the residual update. """ - __slots__ = ("value", "update", "owed") - - def __init__( - self, - value: Optional[torch.Tensor], - update: ResidualUpdate, - owed: Union[UnreducedOutput, DeclaredSum, DeferredFinalize, None] = None, - ): - self.value = value - self.update = update - self.owed = owed + value: Optional[torch.Tensor] + update: ResidualUpdate + owed: Union[UnreducedOutput, DeclaredSum, DeferredFinalize, None] = None def for_boundary(self): # These forms are private inputs to the existing fused-kernel adapters. @@ -81,13 +71,10 @@ def release(self): self.owed = None -class OwedOutput: +class OwedOutput(msgspec.Struct, frozen=True): """Opaque model-facing handle. Only its boundary may read the contribution.""" - __slots__ = ("contribution",) - - def __init__(self, contribution: Contribution): - self.contribution = contribution + contribution: Contribution class ResidualStream: @@ -214,11 +201,8 @@ def snapshot(self, hidden): elif isinstance(pending.owed, DeclaredSum): value = pending.owed.complete(pending.value.clone()) elif isinstance(pending.owed, UnreducedOutput): - owed = pending.owed - value = UnreducedOutput( - pending.value.clone(), - group=owed.group, - reduce_to_dp_local=owed.reduce_to_dp_local, + value = msgspec.structs.replace( + pending.owed, partial=pending.value.clone() ).complete() else: raise NotImplementedError("a finalize handoff requires main-output capture") diff --git a/test/registered/unit/layer_boundary/test_ffn_exit.py b/test/registered/unit/layer_boundary/test_ffn_exit.py index f39fc83211c2..1728c9093207 100644 --- a/test/registered/unit/layer_boundary/test_ffn_exit.py +++ b/test/registered/unit/layer_boundary/test_ffn_exit.py @@ -266,40 +266,6 @@ def test_deferral_implies_fusion_and_passes_the_handoff_through(self): ) self.assertIsInstance(hidden_states, UnreducedOutput) - def test_compiles_without_graph_breaks(self): - """The exit and the output it leaves trace under fullgraph torch.compile.""" - group = types.SimpleNamespace(all_reduce=lambda h: h * 3) - for defers in (False, True): - with self.subTest(defers=defers): - communicator = stub_plan() - communicator.terminal = False - communicator.paths[BatchVariant.ORDINARY] = ordinary_steps( - OutputContract( - Layout(frozenset()), - group=SumGroup.TP, - may_defer_to_next=True, - ) - ) - output = communicator.output - output._defers_sum = lambda fb, steps, **_: defers - output._sum_in_reduce_scatter = lambda steps, dp_step: False - output.ffn_reduction_group = lambda steps: group - output._complete_now = lambda h, r, **_: (h + 1, r) - - def layer(hidden_states, residual): - stream = ResidualStream(residual) - with output.ffn_exit(None, stream=stream) as ffn_exit: - if get_forward().defer_moe_finalize: - hidden_states = hidden_states * 2 - return stream.complete(ffn_exit.finish(hidden_states)) - - torch._dynamo.reset() - compiled = torch.compile(layer, backend="eager", fullgraph=True) - with patch_communicator("_batch_shards_over_cp", lambda fb: False): - hidden_states = compiled(self.hidden_states, self.residual) - expected = self.hidden_states * 3 if defers else self.hidden_states + 1 - torch.testing.assert_close(hidden_states, expected) - class TestReduceOutput(CustomTestCase): def setUp(self): diff --git a/test/registered/unit/layer_boundary/test_residual_stream.py b/test/registered/unit/layer_boundary/test_residual_stream.py index 154eeaf38bf8..08dab8e8be83 100644 --- a/test/registered/unit/layer_boundary/test_residual_stream.py +++ b/test/registered/unit/layer_boundary/test_residual_stream.py @@ -17,7 +17,6 @@ bind_entry, ) from sglang.srt.layers.layer_boundary.contracts import BatchVariant, StageKind -from sglang.srt.layers.layer_boundary.fusions.cutedsl import MoeDeferredFinalize from sglang.srt.layers.layer_boundary.output import UnreducedOutput from sglang.srt.layers.layer_boundary.residual import batch as residual_batch from sglang.srt.layers.layer_boundary.residual.add_norm import PLAIN_ADD @@ -425,45 +424,6 @@ def update_and_read(self, update, value, residual, norm, **kwargs): with self.assertRaises(RuntimeError): stream.input(hidden) - def test_layers_trace_under_fullgraph(self): - """A stream's record, complete, snapshot and write trace under fullgraph.""" - group = SimpleNamespace(all_reduce=lambda x: x * 2) - - def finalize(rows): - return MoeDeferredFinalize( - routed_output=rows, - expert_weights=rows, - permuted_indices=rows, - gated_shared_output=rows, - m=rows.shape[0], - finish=lambda: rows * 5, - ) - - def layers(hidden, residual): - stream = ResidualStream(residual) - for _ in range(2): - owed = stream.record(hidden, PLAIN_ADD, declared_sum=SumGroup.TP) - residual = stream.write(stream.complete(owed) + stream.residual) - owed = stream.record( - UnreducedOutput(residual * 3, group=group), PLAIN_ADD - ) - captured = stream.snapshot(owed) - _, residual = stream.input(owed) - residual = stream.write(stream.complete(owed) + residual) - owed = stream.record(finalize(residual + captured), PLAIN_ADD) - hidden = stream.write(stream.complete(owed) + residual) - return hidden - - with patch( - "sglang.srt.layers.layer_boundary.residual.stream._sum_group", - lambda sum_group: group, - ): - expected = layers(torch.ones(2, 4), torch.full((2, 4), 3.0)) - torch._dynamo.reset() - compiled = torch.compile(layers, fullgraph=True, backend="eager") - actual = compiled(torch.ones(2, 4), torch.full((2, 4), 3.0)) - torch.testing.assert_close(actual, expected) - class TestBatchStageOwnership(CustomTestCase): def test_terminal_norm_releases_layer_buffers(self): diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index 4f56948e5c1b..9147f84b9b26 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -2414,25 +2414,6 @@ def test_the_package_that_defines_them_is_not_warned_at(self): pass self.assertEqual([str(w.message) for w in seen], []) - def test_in_package_all_reduces_trace_under_fullgraph(self): - """The package's own all-reduces trace under fullgraph torch.compile.""" - import torch - - from sglang.srt.distributed import communication_op, parallel_state - - group = SimpleNamespace(all_reduce=lambda x: x * 2) - with ( - patch.object(parallel_state, "_ALREADY_WARNED", set()), - get_parallel().override(tp_group=group, attn_tp_group=group), - ): - for helper in ( - communication_op.tensor_model_parallel_all_reduce, - communication_op.attention_tensor_model_parallel_all_reduce, - ): - with self.subTest(helper.__name__): - compiled = torch.compile(helper, fullgraph=True, backend="eager") - self.assertEqual(compiled(torch.ones(2)).tolist(), [2.0, 2.0]) - def test_the_guard_would_notice_a_caller(self): self.assertTrue(self._callers("get_self_pp_group"))