Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
2 changes: 2 additions & 0 deletions python/sglang/srt/configs/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -2376,6 +2376,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():
Expand Down
14 changes: 9 additions & 5 deletions python/sglang/srt/layers/radix_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -188,18 +188,19 @@ 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")
attn_out = q.new_empty(
(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,
Expand Down Expand Up @@ -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(
Expand Down
59 changes: 59 additions & 0 deletions test/registered/unit/layers/test_radix_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading