Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
27 commits
Select commit Hold shift + click to select a range
4abf816
Synchronize EAGLE DSA graph fallback across DP ranks
weireweire Jul 23, 2026
edb4b06
Preserve draft DP token counts across MoE backends
weireweire Jul 25, 2026
486de18
Cover DSA graph fallback in disagg draft input test
weireweire Jul 29, 2026
d44420b
Avoid dynamic speculative graph compatibility lookup
weireweire Aug 11, 2026
d6504a6
Fix fused TopK disaggregation test model fixture
weireweire Aug 12, 2026
b9e8946
Merge remote-tracking branch 'origin/main' into fix/issue32182-sync-d…
weireweire Aug 21, 2026
247cbab
Fix EAGLE eager DSA padding across DP ranks
weireweire Aug 21, 2026
e6b8977
Merge main and preserve EAGLE DP graph metadata
weireweire Aug 28, 2026
1df123f
Merge main and preserve draft DP metadata semantics
weireweire Sep 1, 2026
5415420
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
nvpohanh Sep 1, 2026
475be56
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
nvpohanh Sep 2, 2026
2d42a6d
Fix EAGLE disaggregation test model fixtures
weireweire Sep 4, 2026
5d29211
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
nvpohanh Sep 4, 2026
99e02a9
Merge main and separate draft graph votes from prefill prefix metadata
weireweire Sep 14, 2026
04c51a7
Use complete MoE backend enum in DSA padding test
weireweire Sep 14, 2026
4532117
Bind TRT-LLM sparse page-table helper in DSA test
weireweire Sep 14, 2026
0a6fa8c
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
nvpohanh Sep 15, 2026
04a54ce
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
nvpohanh Sep 15, 2026
d492674
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
nvpohanh Sep 16, 2026
61c284d
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
nvpohanh Sep 17, 2026
ad8d626
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
nvpohanh Sep 17, 2026
883f63b
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
weireweire Sep 18, 2026
69de5a5
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
nvpohanh Sep 21, 2026
e75906e
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
kpham-sgl Sep 21, 2026
5fbc285
Add is_draft_worker to mock ModelRunner in mamba prefill track test
weireweire Sep 28, 2026
2285468
Merge branch 'main' into fix/issue32182-sync-draft-graph-fallback
weireweire Sep 28, 2026
eb2f424
Fix MixedMoE e2e decode args for current server args
weireweire Sep 28, 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
1 change: 1 addition & 0 deletions python/sglang/srt/batch_overlap/two_batch_overlap.py
Original file line number Diff line number Diff line change
Expand Up @@ -756,6 +756,7 @@ def filter_batch(
"dp_spec_prefill_coordination_applied",
"return_logprob",
"can_run_decode_cuda_graph",
"can_run_dp_draft_cuda_graph",
"can_run_dp_prefill_cuda_graph",
"dp_prefill_cuda_graph_max_prefix_len",
"dp_padding_mode",
Expand Down
32 changes: 28 additions & 4 deletions python/sglang/srt/layers/attention/dsa/dsa_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -201,6 +201,19 @@ def _broadcast_indexer_topk_from_rank0(
return topk_indices


def _make_eager_idle_topk_result(
x: torch.Tensor, index_topk: int, return_indices: bool
) -> Optional[torch.Tensor]:
if not return_indices:
return None
return torch.full(
(x.shape[0], index_topk),
-1,
dtype=torch.int32,
device=x.device,
)


def rotate_activation(x: torch.Tensor) -> torch.Tensor:
if _is_hip:
from fast_hadamard_transform import hadamard_transform
Expand Down Expand Up @@ -1591,6 +1604,21 @@ def forward_cuda(
layer_id: int,
return_indices: bool = True,
) -> Optional[torch.Tensor]:
# A padded eager IDLE rank has no real requests and therefore no valid
# DSA page-table rows to index. Return invalid rows with the physical DP
# shape so later MLP/EP collectives still agree across ranks.
x_meta = x[0] if isinstance(x, tuple) else x
if (
_is_cuda
and forward_batch.forward_mode.is_idle()
and not get_is_capture_mode()
):
topk_result = _make_eager_idle_topk_result(
x_meta, self.index_topk, return_indices
)
topk_result = _broadcast_indexer_topk_from_rank0(topk_result)
return maybe_capture_indexer_topk(layer_id, topk_result)

if _is_hip:
from sglang.kernels.ops.attention.dsa.tilelang_kernel import act_quant
elif not _is_npu:
Expand All @@ -1599,10 +1627,6 @@ def forward_cuda(
if TYPE_CHECKING:
assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool)

# When upstream uses fused FP8 RMSNorm+quant, activations may be passed as
# a tuple like (x_fp8, x_scale[, y]). Use `x_meta` for shape/device queries.
x_meta = x[0] if isinstance(x, tuple) else x

in_piecewise_or_breakable_cuda_graph = (
_is_in_piecewise_or_breakable_cuda_graph()
)
Expand Down
63 changes: 58 additions & 5 deletions python/sglang/srt/layers/attention/dsa_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -173,6 +173,46 @@ def _to_2d_context_lens(seqlens_32: torch.Tensor, batch_size: int) -> torch.Tens
return seqlens_32.contiguous().view(-1, 1)


def _trim_trtllm_decode_dp_padding(
q_all: torch.Tensor,
topk_indices: Optional[torch.Tensor],
real_batch_size: int,
) -> Tuple[torch.Tensor, Optional[torch.Tensor], int]:
"""Align eager decode inputs with metadata planned before DP padding."""
physical_batch_size = q_all.shape[0]
assert real_batch_size <= physical_batch_size, (
f"DSA metadata batch size ({real_batch_size}) exceeds q batch size "
f"({physical_batch_size})"
)
if topk_indices is not None:
assert real_batch_size <= topk_indices.shape[0], (
f"DSA metadata batch size ({real_batch_size}) exceeds topk batch size "
f"({topk_indices.shape[0]})"
)

num_padding_rows = physical_batch_size - real_batch_size
if num_padding_rows == 0:
return q_all, topk_indices, 0

return (
q_all[:real_batch_size],
topk_indices[:real_batch_size] if topk_indices is not None else None,
num_padding_rows,
)


def _restore_trtllm_decode_dp_padding(
output: torch.Tensor, num_padding_rows: int
) -> torch.Tensor:
"""Restore the physical DP shape required by downstream MLP collectives."""
if num_padding_rows == 0:
return output
return torch.cat(
[output, output.new_zeros((num_padding_rows, *output.shape[1:]))],
dim=0,
)


@dataclass(frozen=True)
class DSAFlashMLAMetadata:
"""Metadata only needed by FlashMLA"""
Expand Down Expand Up @@ -3441,9 +3481,24 @@ def _forward_trtllm(
else:
q_all = q.view(-1, layer.tp_q_head_num, layer.head_dim)

# Eager DP attention can pad q beyond metadata that was deliberately
# planned on the real draft batch. Pad top-k to the physical q shape,
# then run decode attention only on metadata-backed rows. The output is
# restored below before downstream MLP/EP collectives.
if (self.use_fused_topk or not is_prefill) and topk_indices is not None:
topk_indices = self._pad_topk_indices(topk_indices, q.shape[0])

num_decode_padding_rows = 0
if not is_prefill:
q_all, topk_indices, num_decode_padding_rows = (
_trim_trtllm_decode_dp_padding(
q_all,
topk_indices,
metadata.cache_seqlens_int32.shape[0],
)
)

if self.use_fused_topk:
if topk_indices is not None:
topk_indices = self._pad_topk_indices(topk_indices, q.shape[0])
page_table_1 = self._get_fused_topk_page_table(topk_indices)
elif is_prefill:
page_table_1 = transform_index_page_table_prefill(
Expand All @@ -3459,8 +3514,6 @@ def _forward_trtllm(
cu_seqlens_q=metadata.cu_seqlens_q,
)
else:
if topk_indices is not None:
topk_indices = self._pad_topk_indices(topk_indices, q.shape[0])
page_table_1 = transform_index_page_table_decode(
page_table=metadata.page_table_1,
topk_indices=topk_indices,
Expand Down Expand Up @@ -3509,7 +3562,7 @@ def _forward_trtllm(
multi_ctas_kv_counter_buffer=multi_ctas_kv_counter_buffer,
)

return out
return _restore_trtllm_decode_dp_padding(out, num_decode_padding_rows)

def _pad_topk_indices(
self, topk_indices: torch.Tensor, num_tokens: int
Expand Down
7 changes: 7 additions & 0 deletions python/sglang/srt/managers/schedule_batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -2471,6 +2471,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# For DP attention
is_extend_in_batch: bool = False
can_run_decode_cuda_graph: bool = False
# Rank-consistent EAGLE draft replay gate. Keep it separate so missing
# draft-only state does not disable target verification or draft extend.
can_run_dp_draft_cuda_graph: bool = False
can_run_dp_prefill_cuda_graph: bool = False
dp_prefill_cuda_graph_max_prefix_len: int = 0
tbo_split_seq_index: Optional[int] = None
Expand Down Expand Up @@ -2524,6 +2527,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# For DP attention
global_num_tokens: Optional[List[int]] = None
global_num_tokens_for_logprob: Optional[List[int]] = None
# The draft model can use a different MoE A2A backend than the target.
draft_global_num_tokens: Optional[List[int]] = None
draft_global_num_tokens_for_logprob: Optional[List[int]] = None
# Full DP token vector retained for Aiter MegaMoE even when the normal MLP
# TP gather path stores only this rank's token count.
global_spec_verify_tier_num_tokens: Optional[List[int]] = None
Expand Down Expand Up @@ -3831,6 +3837,7 @@ def copy(self):
global_num_tokens=self.global_num_tokens,
global_num_tokens_for_logprob=self.global_num_tokens_for_logprob,
can_run_decode_cuda_graph=self.can_run_decode_cuda_graph,
can_run_dp_draft_cuda_graph=self.can_run_dp_draft_cuda_graph,
can_run_dp_prefill_cuda_graph=self.can_run_dp_prefill_cuda_graph,
dp_prefill_cuda_graph_max_prefix_len=self.dp_prefill_cuda_graph_max_prefix_len,
is_extend_in_batch=self.is_extend_in_batch,
Expand Down
53 changes: 53 additions & 0 deletions python/sglang/srt/managers/scheduler_components/dp_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@
get_memory,
get_parallel,
get_schedule,
get_spec,
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils.common import require_mlp_tp_gather
Expand All @@ -45,6 +46,25 @@
_ENABLE_METRICS_DP_ATTENTION = envs.SGLANG_ENABLE_METRICS_DP_ATTENTION.get()


def _spec_input_cuda_graph_compatible(
local_batch: Optional[ScheduleBatch],
) -> bool:
"""Return the local speculative-draft graph admission bit.

None/idle/prebuilt inputs stay permissive so an active rank with complete
runtime state can still use graphs. Any active incompatible input is
min-reduced by ``MLPSyncBatchInfo`` and forces every DP rank eager.
"""
if (
local_batch is None
or local_batch.forward_mode.is_idle()
or local_batch.forward_mode.is_prebuilt()
):
return True
spec_info = local_batch.spec_info
return spec_info is None or spec_info.cuda_graph_compatible


def _resolve_elastic_world_dp_size(
dp_size: int,
*,
Expand Down Expand Up @@ -92,6 +112,7 @@ class MLPSyncBatchInfo:
num_tokens: int
num_tokens_for_logprob: int
can_run_decode_cuda_graph: bool
can_run_draft_cuda_graph: bool
can_run_prefill_cuda_graph: bool
is_extend_in_batch: bool
local_can_run_tbo: bool
Expand All @@ -117,6 +138,7 @@ def _get_local_tensor(self, device, dtype=torch.int64) -> torch.Tensor:
self.local_forward_mode,
int(self.can_run_prefill_cuda_graph),
self.prefill_cuda_graph_max_prefix_len,
int(self.can_run_draft_cuda_graph),
],
device=device,
dtype=dtype,
Expand All @@ -133,6 +155,7 @@ def _get_fallback_tensor(self, device, dtype=torch.int64) -> torch.Tensor:
ForwardMode.IDLE.value, # local_forward_mode
0, # can_run_prefill_cuda_graph
0, # prefill_cuda_graph_max_prefix_len
1, # can_run_draft_cuda_graph
],
device=device,
dtype=dtype,
Expand Down Expand Up @@ -215,6 +238,7 @@ def all_gather(
self.is_extend_in_batch = bool(tp0_info_cpu[:, 3].max())
self.can_run_prefill_cuda_graph = bool(tp0_info_cpu[:, 6].min())
self.prefill_cuda_graph_max_prefix_len = int(tp0_info_cpu[:, 7].max())
self.can_run_draft_cuda_graph = bool(tp0_info_cpu[:, 8].min())
if _ENABLE_METRICS_DP_ATTENTION:
self.dp_cooperation_info = DPCooperationInfo.create(
tp0_info_cpu[:, 5].tolist()
Expand All @@ -225,6 +249,7 @@ def _update_gather_batch(
batch: ScheduleBatch,
mlp_sync_info: MLPSyncBatchInfo,
require_mlp_tp_gather: bool,
draft_require_mlp_tp_gather: Optional[bool] = None,
skip_global_metadata=False,
):
if not require_mlp_tp_gather:
Expand All @@ -235,6 +260,19 @@ def _update_gather_batch(
batch.global_num_tokens_for_logprob = (
mlp_sync_info.global_num_tokens_for_logprob
)
# Reuse the same all-gather result for a draft model whose A2A backend
# requires a different local/full token-count representation.
if draft_require_mlp_tp_gather is not None:
if draft_require_mlp_tp_gather:
batch.draft_global_num_tokens = mlp_sync_info.global_num_tokens

@Fridge003 Fridge003 Aug 21, 2026 •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

(from Codex)

When the draft A2A backend is none, require_mlp_tp_gather() makes this branch use the all-rank counts. For a sufficiently occupied uneven decode batch, DpPaddingMode selects MAX_LEN, so the eager draft ForwardBatch is padded to the largest rank.

The EAGLE eager fallback, however, pre-plans DSA metadata before that padding and marks it ready. An idle rank can therefore have padded query rows with a zero-row page table, while a shorter active rank can reach _forward_trtllm() with more query rows than page_table_1 rows. This can fail in paged-MQA or reshape code and potentially strand peer ranks. Please trim eager DSA query/top-k tensors to the real metadata row count and restore the output padding afterward, short-circuit the eager IDLE indexer, and add mixed active/idle plus uneven-count DP coverage.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This make sense, should be a preexisting issue for dsa to handle padding. added the trim and restore method.

batch.draft_global_num_tokens_for_logprob = (
mlp_sync_info.global_num_tokens_for_logprob
)
else:
batch.draft_global_num_tokens = [mlp_sync_info.num_tokens]
batch.draft_global_num_tokens_for_logprob = [
mlp_sync_info.num_tokens_for_logprob
]
if envs.SGLANG_ENABLE_DP_SPEC_PREFILL_COORDINATION.get():
# Fresh counts have not yet been adjusted by the coordination plan.
batch.dp_spec_prefill_coordination_applied = False
Expand All @@ -251,6 +289,7 @@ def _update_gather_batch(

# Check forward mode for cuda graph
batch.can_run_decode_cuda_graph = mlp_sync_info.can_run_decode_cuda_graph
batch.can_run_dp_draft_cuda_graph = mlp_sync_info.can_run_draft_cuda_graph
batch.can_run_dp_prefill_cuda_graph = mlp_sync_info.can_run_prefill_cuda_graph
batch.dp_prefill_cuda_graph_max_prefix_len = (
mlp_sync_info.prefill_cuda_graph_max_prefix_len
Expand Down Expand Up @@ -378,6 +417,7 @@ def prepare_mlp_sync_batch_raw(
require_mlp_tp_gather: bool,
disable_overlap_schedule: bool,
offload_tags: set[str],
draft_require_mlp_tp_gather: Optional[bool] = None,
dwdp: bool = False,
):
parallel = get_parallel()
Expand Down Expand Up @@ -413,6 +453,7 @@ def prepare_mlp_sync_batch_raw(
can_run_decode_cuda_graph = _local_decode_cuda_graph_vote(
local_batch=local_batch, disable_cuda_graph=disable_cuda_graph
)
can_run_draft_cuda_graph = _spec_input_cuda_graph_compatible(local_batch)
breakable_prefill = check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
full_prefill = check_cuda_graph_backend(Phase.PREFILL, Backend.FULL)
coordinated_prefill = breakable_prefill or full_prefill
Expand Down Expand Up @@ -472,6 +513,7 @@ def prepare_mlp_sync_batch_raw(
num_tokens=num_tokens,
num_tokens_for_logprob=num_tokens_for_logprob,
can_run_decode_cuda_graph=can_run_decode_cuda_graph,
can_run_draft_cuda_graph=can_run_draft_cuda_graph,
can_run_prefill_cuda_graph=can_run_prefill_cuda_graph,
is_extend_in_batch=is_extend_in_batch,
local_can_run_tbo=local_can_run_tbo,
Expand Down Expand Up @@ -515,6 +557,7 @@ def prepare_mlp_sync_batch_raw(
batch_to_gather,
mlp_sync_info,
require_mlp_tp_gather,
draft_require_mlp_tp_gather,
skip_global_metadata=not metadata_ready,
)

Expand Down Expand Up @@ -546,12 +589,22 @@ class SchedulerDPAttnAdapter:
get_require_mlp_sync: Callable[[], bool]

def prepare_mlp_sync_batch(self, local_batch: ScheduleBatch):
draft_require_mlp_tp_gather = None
if self.spec_algorithm.is_eagle() or self.spec_algorithm.is_standalone():
draft_moe_a2a_backend = get_spec().speculative_moe_a2a_backend
if draft_moe_a2a_backend is None:
draft_moe_a2a_backend = get_exec().moe.moe_a2a_backend
draft_require_mlp_tp_gather = require_mlp_tp_gather(
moe_a2a_backend=draft_moe_a2a_backend,
)

return prepare_mlp_sync_batch_raw(
local_batch,
model_runner=self.model_runner,
get_idle_batch=self.get_idle_batch,
disable_cuda_graph=cuda_graph_fully_disabled(),
require_mlp_tp_gather=require_mlp_tp_gather(),
draft_require_mlp_tp_gather=draft_require_mlp_tp_gather,
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
offload_tags=self.offload_tags,
dwdp=get_parallel().dwdp_size > 1,
Expand Down
Loading
Loading