diff --git a/python/sglang/srt/debug_utils/pr_fix_toggle.py b/python/sglang/srt/debug_utils/pr_fix_toggle.py index aecdca3e1..4f7e06185 100644 --- a/python/sglang/srt/debug_utils/pr_fix_toggle.py +++ b/python/sglang/srt/debug_utils/pr_fix_toggle.py @@ -27,12 +27,8 @@ patches: - target: sglang.srt.speculative.eagle_draft_cuda_graph_runner.EAGLEDraftCudaGraphRunner.capture_one_shape edits: - match: | - forward_batch.spec_info.hidden_states = hidden_states_backup forward_batch.positions.sub_(self.eagle_worker.speculative_num_steps - 1) - return ret - replacement: | - forward_batch.spec_info.hidden_states = hidden_states_backup - return ret + replacement: "" """ diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index f0860135d..bfdf49096 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -1640,6 +1640,7 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin): output_topk_p, output_topk_index, output_hidden_states, + output_dsa_topk_indices, output_bootstrap_room, ) = self.metadata_buffers.get_buf(idx) @@ -1733,6 +1734,12 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin): decode_req.req.output_topk_p = output_topk_p decode_req.req.output_topk_index = output_topk_index decode_req.req.hidden_states_tensor = output_hidden_states + if ( + output_dsa_topk_indices is not None + and torch.all(output_dsa_topk_indices < 0).item() + ): + output_dsa_topk_indices = None + decode_req.req.output_dsa_topk_indices = output_dsa_topk_indices if decode_req.req.return_logprob and not replayed_boundary: decode_req.req.logprob.output_token_logprobs_val.append( diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 2f9931644..1da2a091d 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -657,8 +657,14 @@ class SchedulerDisaggregationPrefillMixin: req.hidden_states_tensor = ( batch.spec_info.hidden_states[i].cpu().clone() ) + dsa_topk_indices = batch.spec_info.dsa_topk_indices + if dsa_topk_indices is not None: + req.output_dsa_topk_indices = dsa_topk_indices[i].cpu().clone() + else: + req.output_dsa_topk_indices = None else: req.hidden_states_tensor = None + req.output_dsa_topk_indices = None if req.return_logprob: assert extend_logprob_start_len_per_req is not None assert extend_input_len_per_req is not None @@ -1158,6 +1164,7 @@ class SchedulerDisaggregationPrefillMixin: req.start_send_idx = 0 req.tmp_end_idx = -1 req.hidden_states_tensor = None + req.output_dsa_topk_indices = None req.pending_bootstrap = True req.time_stats.reset_prefill_retry_time() if req.prefill_attempt_count >= max_attempts: diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index a000cf5fb..755d1e00c 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -11,6 +11,7 @@ import numpy as np import torch import torch.distributed as dist +from sglang.srt.configs.model_config import get_dsa_index_topk from sglang.srt.disaggregation.base import KVPoll from sglang.srt.environ import envs from sglang.srt.utils import is_hip, is_npu @@ -33,6 +34,13 @@ FAKE_BOOTSTRAP_HOST = "2.2.2.2" _IS_HIP = is_hip() +def get_dsa_seed_metadata_dim(hf_config) -> int: + """Return the model-defined PD seed width, independent of local spec mode.""" + if not getattr(hf_config, "index_share_for_mtp_iteration", False): + return 0 + return get_dsa_index_topk(hf_config) + + def is_dsv4_c128_online_enabled() -> bool: """Return whether DSV4 C128 uses request-scoped online state.""" return not _IS_HIP and envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get() @@ -228,8 +236,10 @@ class MetadataBuffers: max_top_logprobs_num: int = 128, max_sampling_mask_tokens: Optional[int] = None, custom_mem_pool: torch.cuda.MemPool = None, + output_dsa_topk_indices_dim: int = 0, ): self.custom_mem_pool = custom_mem_pool + self.output_dsa_topk_indices_dim = output_dsa_topk_indices_dim if max_sampling_mask_tokens is None: max_sampling_mask_tokens = ( envs.SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS.get() @@ -295,6 +305,15 @@ class MetadataBuffers: self.output_hidden_states = torch.zeros( (size, hidden_size), dtype=hidden_states_dtype, device=device ) + if self.output_dsa_topk_indices_dim > 0: + self.output_dsa_topk_indices = torch.full( + (size, self.output_dsa_topk_indices_dim), + -1, + dtype=torch.int32, + device=device, + ) + else: + self.output_dsa_topk_indices = None # Request validation: store bootstrap_room to detect metadata corruption self.bootstrap_room = torch.zeros( (size, 8), dtype=bootstrap_room_dtype, device=device @@ -322,9 +341,11 @@ class MetadataBuffers: self.output_topk_p, self.output_topk_index, self.output_hidden_states, - self.bootstrap_room, ] ) + if self.output_dsa_topk_indices is not None: + bufs.append(self.output_dsa_topk_indices) + bufs.append(self.bootstrap_room) ptrs = [buf.data_ptr() for buf in bufs] data_lens = [buf.nbytes for buf in bufs] item_lens = [buf[0].nbytes for buf in bufs] @@ -351,6 +372,11 @@ class MetadataBuffers: self.output_topk_p[idx].clone(), self.output_topk_index[idx].clone(), self.output_hidden_states[idx].clone(), + ( + self.output_dsa_topk_indices[idx].clone() + if self.output_dsa_topk_indices is not None + else None + ), self.bootstrap_room[idx].clone(), ) @@ -461,6 +487,14 @@ class MetadataBuffers: self.output_hidden_states[req.metadata_buffer_index].copy_( req.hidden_states_tensor ) + if self.output_dsa_topk_indices is not None: + dsa_topk_indices = req.output_dsa_topk_indices + if dsa_topk_indices is not None: + self.output_dsa_topk_indices[req.metadata_buffer_index].copy_( + dsa_topk_indices + ) + else: + self.output_dsa_topk_indices[req.metadata_buffer_index].fill_(-1) # Store bootstrap_room for validation on decode side self.bootstrap_room[req.metadata_buffer_index, 0] = ( req.bootstrap_room if req.bootstrap_room is not None else 0 diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index 21a684541..e15c822db 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -67,6 +67,17 @@ def compute_dsa_seqlens(original_seq_lens, dsa_index_topk: int): return original_seq_lens.clamp(max=dsa_index_topk) +def should_use_dsa_fused_topk( + server_args, seed_dsa_topk_from_draft_extend: bool +) -> bool: + pd_index_share_seed = ( + server_args.disaggregation_mode != "null" and seed_dsa_topk_from_draft_extend + ) + # TODO(kpham-sgl): Transfer request-relative IndexShare seeds and remap them + # to decode-local KV slots so fused top-k can remain enabled under PD. + return envs.SGLANG_DSA_FUSE_TOPK.get() and not pd_index_share_seed + + def is_dsa_enable_prefill_cp(): return get_server_args().enable_dsa_prefill_context_parallel diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index f3e637d6c..c716b0ea3 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -45,6 +45,7 @@ from sglang.srt.layers.attention.dsa.utils import ( is_dsa_enable_prefill_cp, is_dsa_prefill_cp_in_seq_split, pad_dsa_cache_seqlens, + should_use_dsa_fused_topk, ) from sglang.srt.layers.attention.utils import ( concat_mla_absorb_q_general, @@ -63,6 +64,7 @@ from sglang.srt.utils import ( is_gfx95_supported, is_hip, is_sm100_supported, + print_warning_once, ) # Opt-in (default off): route the fp8 sparse-MLA prefill path through the Triton @@ -340,6 +342,7 @@ class DeepseekSparseAttnBackend( speculative_step_id=0, topk=0, speculative_num_steps=0, + seed_dsa_topk_from_draft_extend: bool = False, ): super().__init__() self.forward_metadata: DSAMetadata @@ -433,6 +436,13 @@ class DeepseekSparseAttnBackend( model_runner.server_args.speculative_num_draft_tokens ) self.speculative_step_id = speculative_step_id + self.use_fused_topk = should_use_dsa_fused_topk( + model_runner.server_args, seed_dsa_topk_from_draft_extend + ) + if envs.SGLANG_DSA_FUSE_TOPK.get() and not self.use_fused_topk: + print_warning_once( + "Disabling fused DSA top-k for IndexShare under PD disaggregation." + ) self.device_capability = torch.cuda.get_device_capability() self.device_sm_major = self.device_capability[0] @@ -1107,7 +1117,7 @@ class DeepseekSparseAttnBackend( and self.real_page_size > 1 and self.hisparse_coordinator is None and not self.speculative_num_draft_tokens - and envs.SGLANG_DSA_FUSE_TOPK.get() + and self.use_fused_topk and envs.SGLANG_OPT_USE_TOPK_V2.get() and self.dsa_index_topk is not None and self.dsa_index_topk <= 2048 @@ -1897,7 +1907,7 @@ class DeepseekSparseAttnBackend( topk_transform_method = self.get_topk_transform_method( forward_batch.forward_mode ) - if envs.SGLANG_DSA_FUSE_TOPK.get(): + if self.use_fused_topk: page_table_1 = self._get_fused_topk_page_table(topk_indices) else: if topk_transform_method == TopkTransformMethod.RAGGED: @@ -2112,7 +2122,7 @@ class DeepseekSparseAttnBackend( topk_indices, layer.layer_id, ) - elif envs.SGLANG_DSA_FUSE_TOPK.get(): + elif self.use_fused_topk: page_table_1 = self._get_fused_topk_page_table(topk_indices) else: page_table_1 = transform_index_page_table_decode( @@ -2669,7 +2679,7 @@ class DeepseekSparseAttnBackend( if topk_indices is not None: topk_indices = self._pad_topk_indices(topk_indices, q.shape[0]) - if envs.SGLANG_DSA_FUSE_TOPK.get(): + if self.use_fused_topk: page_table_1 = self._get_fused_topk_page_table(topk_indices) elif is_prefill: page_table_1 = transform_index_page_table_prefill( @@ -2835,7 +2845,7 @@ class DeepseekSparseAttnBackend( def get_indexer_metadata( self, layer_id: int, forward_batch: ForwardBatch ) -> DSAIndexerMetadata: - force_unfused = ( + force_unfused = not self.use_fused_topk or ( self.hisparse_coordinator is not None and forward_batch.forward_mode.is_decode_or_idle() ) @@ -2879,7 +2889,11 @@ class DeepseekSparseAttnMultiStepBackend: needs_cpu_seq_lens: bool = False def __init__( - self, model_runner: ModelRunner, topk: int, speculative_num_steps: int + self, + model_runner: ModelRunner, + topk: int, + speculative_num_steps: int, + seed_dsa_topk_from_draft_extend: bool = False, ): self.topk = topk self.speculative_num_steps = speculative_num_steps @@ -2891,6 +2905,7 @@ class DeepseekSparseAttnMultiStepBackend: speculative_step_id=i, topk=self.topk, speculative_num_steps=self.speculative_num_steps, + seed_dsa_topk_from_draft_extend=seed_dsa_topk_from_draft_extend, ) ) diff --git a/python/sglang/srt/managers/overlap_utils.py b/python/sglang/srt/managers/overlap_utils.py index 5777ac983..9e2f191aa 100644 --- a/python/sglang/srt/managers/overlap_utils.py +++ b/python/sglang/srt/managers/overlap_utils.py @@ -146,7 +146,7 @@ class RelayPayload: topk_index=draft_input.topk_index, hidden_states=draft_input.hidden_states, draft_probs=getattr(draft_input, "draft_probs", None), - dsa_topk_indices=getattr(draft_input, "dsa_topk_indices", None), + dsa_topk_indices=draft_input.dsa_topk_indices, ) @@ -281,6 +281,7 @@ class FutureMap: self.fwd_prepare_d2h_stream = None # Lazy-inited on the first non-empty stash (peeks tensor shapes); non-spec's is a no-op. self._forward_buf_initialized = False + self.dsa_topk_indices_buf = None self.publish_ready = None # lazy device.Event(); only spec_v2 needs it # Debug consume-once state: armed by a recording publish, consumed by @@ -338,14 +339,15 @@ class FutureMap: device=self.device, ) - self.dsa_topk_indices_buf = None - if payload.dsa_topk_indices is not None: - seed0 = payload.dsa_topk_indices[0] - self.dsa_topk_indices_buf = torch.empty( - (self.req_pool_size, *seed0.shape), - dtype=payload.dsa_topk_indices.dtype, - device=self.device, - ) + def _maybe_init_dsa_topk_indices_buf(self, payload: RelayPayload) -> None: + if self.dsa_topk_indices_buf is not None or payload.dsa_topk_indices is None: + return + seed0 = payload.dsa_topk_indices[0] + self.dsa_topk_indices_buf = torch.empty( + (self.req_pool_size, *seed0.shape), + dtype=payload.dsa_topk_indices.dtype, + device=self.device, + ) def resolve_confidence_cpu( self, batch: ScheduleBatch @@ -397,8 +399,11 @@ class FutureMap: draft_input.bonus_tokens = self.output_tokens_buf[indices] if self.need_hidden_states and not self.need_topk: draft_input.hidden_states = self.hidden_states_buf[indices] - if self.dsa_topk_indices_buf is not None: + if draft_input.future_dsa_topk_indices_available: + assert self.dsa_topk_indices_buf is not None draft_input.dsa_topk_indices = self.dsa_topk_indices_buf[indices] + else: + draft_input.dsa_topk_indices = None if _DEBUG_ASSERT: _assert_nonneg_and_invalidate( draft_input.bonus_tokens, self.output_tokens_buf, indices @@ -503,6 +508,7 @@ class FutureMap: return if not self._forward_buf_initialized: self._lazy_init_forward_buf(payload) + self._maybe_init_dsa_topk_indices_buf(payload) self.output_tokens_buf[indices] = payload.bonus_tokens.to( self.output_tokens_buf.dtype ) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index e5c59ac0b..c68a235e1 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -933,6 +933,7 @@ class Req(ReqDllmMixin): self.hidden_states_tensor = None # Note: use tensor instead of list to transfer hidden_states when PD + MTP self.output_topk_p = None self.output_topk_index = None + self.output_dsa_topk_indices = None # capture routed experts self.return_routed_experts = return_routed_experts diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index a773dfc33..9f8fc168d 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -61,6 +61,7 @@ from sglang.srt.disaggregation.utils import ( MetadataBuffers, ReqToMetadataIdxAllocator, TransferBackend, + get_dsa_seed_metadata_dim, prepare_abort, ) from sglang.srt.distributed import get_pp_group, get_world_group @@ -72,9 +73,7 @@ from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_r from sglang.srt.layers.attention.mamba.ops import ( initialize_mamba_selective_state_update_backend, ) -from sglang.srt.layers.dp_attention import ( - compute_dp_attention_world_info, -) +from sglang.srt.layers.dp_attention import compute_dp_attention_world_info from sglang.srt.layers.moe import initialize_moe_config from sglang.srt.layers.quantization.fp4_utils import initialize_fp4_gemm_config from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config @@ -1139,6 +1138,12 @@ class Scheduler( disagg_hidden_size = 16 # minimal padding size for RDMA disagg_hidden_states_dtype = torch.float32 + # The PD metadata wire schema must match on P and D even when only D + # enables spec decoding; a seedless prefill writes the invalid sentinel. + output_dsa_topk_indices_dim = get_dsa_seed_metadata_dim( + self.model_config.hf_config + ) + if ( self.disaggregation_mode == DisaggregationMode.DECODE ): # *8 headroom for MiniMax-M3; *2 for other models. @@ -1154,6 +1159,7 @@ class Scheduler( hidden_size=disagg_hidden_size, hidden_states_dtype=disagg_hidden_states_dtype, custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(), + output_dsa_topk_indices_dim=output_dsa_topk_indices_dim, ) # The decode requests polling kv cache @@ -1199,6 +1205,7 @@ class Scheduler( hidden_size=disagg_hidden_size, hidden_states_dtype=disagg_hidden_states_dtype, custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(), + output_dsa_topk_indices_dim=output_dsa_topk_indices_dim, ) self.disagg_prefill_bootstrap_queue = PrefillBootstrapQueue( @@ -3390,6 +3397,9 @@ class Scheduler( if not batch.spec_algorithm.is_none(): batch.spec_info = batch_result.next_draft_input + batch.spec_info.future_dsa_topk_indices_available = ( + batch.spec_info.dsa_topk_indices is not None + ) batch.spec_info.future_indices = future_indices elif self.enable_pdmux and batch.forward_mode.is_split_prefill(): resolve_forward_inputs(batch, self.future_map) diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py index 6f2e3747b..0673ea5be 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py @@ -61,6 +61,36 @@ def _resolve_attn_backend(forward_batch: ForwardBatch): return backend +def _forward_dsa_indexer_for_mha( + indexer, + *, + hidden_states: torch.Tensor, + q_lora: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + layer_id: int, +) -> None: + """Fill the indexer K cache and publish an MTP seed when requested.""" + spec_info = forward_batch.spec_info + seed_buf = spec_info.dsa_seed_topk_capture if spec_info is not None else None + topk_indices = indexer( + x=hidden_states, + q_lora=q_lora, + positions=positions, + forward_batch=forward_batch, + layer_id=layer_id, + return_indices=seed_buf is not None, + ) + if seed_buf is None: + return + if topk_indices is None: + raise RuntimeError("DSA MHA indexer did not produce the requested MTP seed") + + select = spec_info.dsa_seed_topk_select + src = topk_indices if select is None else topk_indices[select] + seed_buf[: src.shape[0]].copy_(src) + + # Configs for DeepSeek-V3: # num_local_heads = 128 # qk_nope_head_dim = 128 @@ -174,13 +204,13 @@ class DeepseekMHAForwardMixin: -1, self.num_local_heads, self.qk_head_dim ) if self.should_run_indexer(): - _ = self.indexer( - x=hidden_states, + _forward_dsa_indexer_for_mha( + self.indexer, + hidden_states=hidden_states, q_lora=q_lora, positions=positions, forward_batch=forward_batch, layer_id=self.layer_id, - return_indices=False, ) elif _use_aiter_gfx95 and self.q_b_proj.weight.dtype == torch.uint8: # MXFP4: fused RMSNorm + quant diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index b831da232..5c254e9a2 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -20,11 +20,13 @@ class DraftBackendFactory: draft_model_runner, topk: int, speculative_num_steps: int, + seed_dsa_topk_from_draft_extend: bool = False, ): self.server_args = server_args self.draft_model_runner = draft_model_runner self.topk = topk self.speculative_num_steps = speculative_num_steps + self.seed_dsa_topk_from_draft_extend = seed_dsa_topk_from_draft_extend self.draft_attn_backend = server_args.speculative_draft_attention_backend def _create_backend( @@ -110,13 +112,20 @@ class DraftBackendFactory: ) return DeepseekSparseAttnMultiStepBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + self.draft_model_runner, + self.topk, + self.speculative_num_steps, + seed_dsa_topk_from_draft_extend=self.seed_dsa_topk_from_draft_extend, ) def _create_dsa_prefill_backend(self): from sglang.srt.layers.attention.dsa_backend import DeepseekSparseAttnBackend - return DeepseekSparseAttnBackend(self.draft_model_runner, skip_prefill=False) + return DeepseekSparseAttnBackend( + self.draft_model_runner, + skip_prefill=False, + seed_dsa_topk_from_draft_extend=self.seed_dsa_topk_from_draft_extend, + ) def _create_flashinfer_decode_backend(self): if not self.draft_model_runner.use_mla_backend: diff --git a/python/sglang/srt/speculative/eagle_disaggregation.py b/python/sglang/srt/speculative/eagle_disaggregation.py index 5a47608a5..57742557c 100644 --- a/python/sglang/srt/speculative/eagle_disaggregation.py +++ b/python/sglang/srt/speculative/eagle_disaggregation.py @@ -51,15 +51,24 @@ def build_eagle_disagg_draft_input( [req.hidden_states_tensor for req in batch.reqs], dim=0 ).to(batch.device) + dsa_topk_indices = None + dsa_indices_list = [req.output_dsa_topk_indices for req in batch.reqs] + if dsa_indices_list and all(t is not None for t in dsa_indices_list): + dsa_topk_indices = torch.stack(dsa_indices_list, dim=0).to(batch.device) + if torch.any(torch.all(dsa_topk_indices < 0, dim=1)).item(): + dsa_topk_indices = None + spec_info = EagleDraftInput( topk_p=topk_p, topk_index=topk_index, hidden_states=hidden_states, bonus_tokens=last_tokens_tensor, + dsa_topk_indices=dsa_topk_indices, ) spec_info.capture_hidden_mode = CaptureHiddenMode.LAST if batch.enable_overlap: + spec_info.future_dsa_topk_indices_available = dsa_topk_indices is not None spec_info.future_indices = batch.req_pool_indices # Seed the relay buf with the known seq_lens; publish's chained record # keeps the in-flight forward's fence intact (see FutureMap.publish). diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index 43f6aefed..9478fb909 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -442,11 +442,13 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): output_cache_loc_backup = forward_batch.out_cache_loc hidden_states_backup = forward_batch.spec_info.hidden_states + dsa_topk_indices_backup = forward_batch.spec_info.dsa_topk_indices ret = self.eagle_worker.draft_forward(forward_batch) forward_batch.out_cache_loc = output_cache_loc_backup forward_batch.spec_info.hidden_states = hidden_states_backup + forward_batch.spec_info.dsa_topk_indices = dsa_topk_indices_backup forward_batch.positions.sub_(self.eagle_worker.speculative_num_steps - 1) return ret @@ -647,6 +649,8 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): ) with timer_ctx: out = self._replay_graph(shape_key, forward_batch) + if self.buffers.dsa_seed_topk is not None: + forward_batch.spec_info.dsa_topk_indices = None if bs != raw_bs: out = self._postprocess_output_to_raw_bs(out, raw_bs) diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index ecdb47880..ca4bd35f0 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -178,6 +178,7 @@ class EagleDraftInput(SpecInput): # V2 overlap worker only: req_pool_indices used as buf slot keys. future_indices: Optional[torch.Tensor] = None + future_dsa_topk_indices_available: bool = False def __post_init__(self): super().__init__(SpecInputType.EAGLE_DRAFT) @@ -255,6 +256,10 @@ class EagleDraftInput(SpecInput): self.future_indices = torch.cat( [self.future_indices, spec_info.future_indices] ) + self.future_dsa_topk_indices_available = ( + self.future_dsa_topk_indices_available + and spec_info.future_dsa_topk_indices_available + ) return # Detect idle stub by `topk_index` length (idle inputs have diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index f0fde3aaa..c5b2ff109 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -373,6 +373,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): self.draft_runner, self.topk, self.speculative_num_steps, + seed_dsa_topk_from_draft_extend=self.seed_dsa_topk_from_draft_extend, ) # Initialize decode attention backend @@ -519,6 +520,13 @@ class EagleDraftWorker(EagleDraftWorkerBase): self.topk, self.speculative_num_steps, ) + if ( + can_cuda_graph + and not forward_batch.forward_mode.is_idle() + and self.seed_dsa_topk_from_draft_extend + and draft_input.dsa_topk_indices is None + ): + can_cuda_graph = False n_inner = self.speculative_num_steps - 1 canary_outside_ctx = ( diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index 96193a608..1e8f1988d 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -281,9 +281,7 @@ class SpeculativeAlgorithm(Enum): return EAGLEWorkerV2 elif self.is_standalone(): - from sglang.srt.speculative.standalone_worker_v2 import ( - StandaloneWorkerV2, - ) + from sglang.srt.speculative.standalone_worker_v2 import StandaloneWorkerV2 return StandaloneWorkerV2 elif self.is_ngram(): @@ -316,6 +314,13 @@ class SpecInput(ABC): # assignment, so an init-time default would clobber the passed layout. ragged_verify_layout: Optional[RaggedVerifyLayout] = None + # DSA MTP IndexShare seed relay. Class-level defaults (same rationale as + # ragged_verify_layout) so scheduler/relay/attention code reads them + # uniformly on any SpecInput; only the EAGLE-family inputs override them. + dsa_topk_indices: Optional[torch.Tensor] = None + future_dsa_topk_indices_available: bool = False + dsa_seed_topk_capture: Optional[torch.Tensor] = None + def __init__(self, spec_input_type: SpecInputType): self.spec_input_type = spec_input_type diff --git a/test/registered/kernels/test_dsa_indexer.py b/test/registered/kernels/test_dsa_indexer.py index 3102f2126..ba841c19c 100644 --- a/test/registered/kernels/test_dsa_indexer.py +++ b/test/registered/kernels/test_dsa_indexer.py @@ -261,6 +261,7 @@ class MockModelRunner: "dsa_decode_backend": "fa3", "dsa_topk_backend": "sgl-kernel", "dsa_paged_mqa_logits_backend": "auto", + "disaggregation_mode": "null", }, )() self.hisparse_coordinator = None diff --git a/test/registered/unit/disaggregation/test_disaggregation_wire.py b/test/registered/unit/disaggregation/test_disaggregation_wire.py index 1540f10b5..b96982b0c 100644 --- a/test/registered/unit/disaggregation/test_disaggregation_wire.py +++ b/test/registered/unit/disaggregation/test_disaggregation_wire.py @@ -13,10 +13,15 @@ from sglang.srt.disaggregation.common.utils import ( unpack_list_of_buffers, ) from sglang.srt.disaggregation.utils import ( + MetadataBuffers, get_dsv4_c128_state_indices, setup_state_kv_args, ) +from sglang.srt.managers.overlap_utils import FutureMap, RelayPayload from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool +from sglang.srt.speculative.eagle_disaggregation import ( + build_eagle_disagg_draft_input, +) from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=2, suite="base-a-test-cpu") @@ -93,6 +98,97 @@ class TestGroupConcurrentContiguous(unittest.TestCase): group_concurrent_contiguous(self._arr([1, 2, 3]), self._arr([1, 2])) +class TestEagleDsaSeedTransfer(unittest.TestCase): + @staticmethod + def _make_req(seed, metadata_buffer_index=0): + return SimpleNamespace( + metadata_buffer_index=metadata_buffer_index, + output_ids=[101], + cached_tokens=0, + cached_tokens_device=0, + cached_tokens_host=0, + cached_tokens_storage=0, + multimodal_inputs=None, + return_logprob=False, + return_sampling_mask=False, + hidden_states_tensor=torch.tensor([1.0, 2.0]), + output_topk_p=torch.tensor([1.0]), + output_topk_index=torch.tensor([7]), + output_dsa_topk_indices=seed, + bootstrap_room=9, + ) + + def test_metadata_buffer_copies_seed_and_uses_invalid_sentinel(self): + buffers = MetadataBuffers( + size=2, + hidden_size=2, + hidden_states_dtype=torch.float32, + output_dsa_topk_indices_dim=3, + ) + seed = torch.tensor([4, 5, 6], dtype=torch.int32) + buffers.set_buf(self._make_req(seed)) + buffers.set_buf(self._make_req(None, metadata_buffer_index=1)) + + self.assertTrue(torch.equal(buffers.output_dsa_topk_indices[0], seed)) + self.assertEqual(buffers.output_dsa_topk_indices[1].tolist(), [-1, -1, -1]) + ptrs, data_lens, item_lens = buffers.get_buf_infos() + self.assertEqual(ptrs[-2], buffers.output_dsa_topk_indices.data_ptr()) + self.assertEqual(data_lens[-2], buffers.output_dsa_topk_indices.nbytes) + self.assertEqual(item_lens[-2], buffers.output_dsa_topk_indices[0].nbytes) + + def test_decode_input_requires_valid_seed_for_every_request(self): + seeds = ( + torch.tensor([1, 2, 3], dtype=torch.int32), + torch.tensor([4, 5, 6], dtype=torch.int32), + ) + batch = SimpleNamespace( + reqs=[self._make_req(seed) for seed in seeds], + device="cpu", + enable_overlap=False, + ) + server_args = SimpleNamespace( + speculative_eagle_topk=1, + speculative_num_steps=5, + enable_multi_layer_eagle=False, + ) + last_tokens = torch.tensor([11, 12], dtype=torch.int64) + + draft_input = build_eagle_disagg_draft_input( + batch, server_args, last_tokens, None + ) + self.assertTrue(torch.equal(draft_input.dsa_topk_indices, torch.stack(seeds))) + + for invalid_seed in ( + None, + torch.full((3,), -1, dtype=torch.int32), + ): + batch.reqs[1].output_dsa_topk_indices = invalid_seed + draft_input = build_eagle_disagg_draft_input( + batch, server_args, last_tokens, None + ) + self.assertIsNone(draft_input.dsa_topk_indices) + + def test_future_map_initializes_seed_buffer_after_seedless_payload(self): + future_map = object.__new__(FutureMap) + future_map.dsa_topk_indices_buf = None + future_map.req_pool_size = 4 + future_map.device = "cpu" + future_map._maybe_init_dsa_topk_indices_buf( + RelayPayload(bonus_tokens=torch.zeros((2,), dtype=torch.int64)) + ) + self.assertIsNone(future_map.dsa_topk_indices_buf) + + seeds = torch.tensor([[1, 2, 3], [4, 5, 6]], dtype=torch.int32) + future_map._maybe_init_dsa_topk_indices_buf( + RelayPayload( + bonus_tokens=torch.zeros((2,), dtype=torch.int64), + dsa_topk_indices=seeds, + ) + ) + self.assertEqual(future_map.dsa_topk_indices_buf.shape, (4, 3)) + self.assertEqual(future_map.dsa_topk_indices_buf.dtype, torch.int32) + + class TestDSV4C128StateIndices(unittest.TestCase): def test_online_aligned_boundary_has_no_partial_state(self): np.testing.assert_array_equal( diff --git a/test/registered/unit/spec/test_eagle_worker_v2_topk1_fastpath.py b/test/registered/unit/spec/test_eagle_worker_v2_topk1_fastpath.py index c7bae0a9a..a93dfb5c6 100644 --- a/test/registered/unit/spec/test_eagle_worker_v2_topk1_fastpath.py +++ b/test/registered/unit/spec/test_eagle_worker_v2_topk1_fastpath.py @@ -8,10 +8,11 @@ slow path (`organize_draft_results`) for num_steps in {1, 2, 3, 4}. import unittest from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import MagicMock, patch import torch +from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.speculative.adaptive_runtime_state import SpecRuntimeState from sglang.srt.speculative.eagle_utils import organize_draft_results from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2 @@ -71,10 +72,11 @@ def _make_worker(num_steps: int, num_draft_tokens: int): return worker -def _make_backend_factory(decode_backend, draft_extend_backend): +def _make_backend_factory(decode_backend, draft_extend_backend, captured_kwargs=None): class FakeDraftBackendFactory: def __init__(self, *args, **kwargs): - pass + if captured_kwargs is not None: + captured_kwargs.update(kwargs) def create_decode_backend(self): return decode_backend @@ -122,6 +124,78 @@ class TestEagleWorkerV2Topk1FastPath(CustomTestCase): class TestEagleWorkerV2BackendFallback(CustomTestCase): + def test_missing_seed_cuda_graph_fallback(self): + graph_result = ( + [], + torch.zeros((1, 1), dtype=torch.long, device=DEVICE), + torch.zeros((1, 1), dtype=torch.long, device=DEVICE), + None, + ) + tree_result = ( + torch.empty((0,), dtype=torch.bool, device=DEVICE), + torch.zeros((1,), dtype=torch.long, device=DEVICE), + torch.zeros((1, 2), dtype=torch.long, device=DEVICE), + torch.zeros((1, 2), dtype=torch.long, device=DEVICE), + torch.zeros((1, 2), dtype=torch.long, device=DEVICE), + torch.zeros((2,), dtype=torch.long, device=DEVICE), + ) + + for seed_enabled, seed_present, expect_graph in ( + (True, False, False), + (True, True, True), + (False, False, True), + ): + with self.subTest( + seed_enabled=seed_enabled, + seed_present=seed_present, + ): + worker = object.__new__(EagleDraftWorker) + worker.req_to_token_pool = None + worker.cuda_graph_runner = SimpleNamespace( + execute=MagicMock(return_value=graph_result) + ) + worker.draft_runner = SimpleNamespace(canary_manager=None) + worker.topk = 1 + worker.speculative_num_steps = 1 + worker.speculative_num_draft_tokens = 2 + worker.device = DEVICE + worker.tree_mask_mode = None + worker.seed_dsa_topk_from_draft_extend = seed_enabled + worker.index_share_for_mtp_iteration = True + forward_batch = SimpleNamespace(forward_mode=ForwardMode.DECODE) + worker.prepare_for_draft = MagicMock(return_value=(forward_batch, True)) + worker.draft_forward = MagicMock(return_value=graph_result) + attn_backend = SimpleNamespace( + get_verify_buffers_to_fill_after_draft=lambda: (None, None), + max_context_len=1, + ) + worker.target_worker = SimpleNamespace( + model_runner=SimpleNamespace(attn_backend=attn_backend) + ) + draft_input = SimpleNamespace( + bonus_tokens=torch.zeros((1,), dtype=torch.long, device=DEVICE), + dsa_topk_indices=( + torch.ones((1, 1), dtype=torch.int32, device=DEVICE) + if seed_present + else None + ), + ) + batch = SimpleNamespace( + spec_info=draft_input, + forward_mode=ForwardMode.DECODE, + seq_lens_sum=1, + seq_lens=torch.ones((1,), dtype=torch.int32, device=DEVICE), + ) + + with patch( + "sglang.srt.speculative.eagle_worker_v2.build_tree_kernel_efficient", + return_value=tree_result, + ): + worker.draft(batch) + + self.assertEqual(worker.cuda_graph_runner.execute.called, expect_graph) + self.assertEqual(worker.draft_forward.called, not expect_graph) + def test_preserves_initialized_backend_when_draft_extend_backend_is_unset(self): worker = object.__new__(EagleDraftWorker) existing_backend = object() @@ -130,6 +204,7 @@ class TestEagleWorkerV2BackendFallback(CustomTestCase): worker.draft_runner = SimpleNamespace(attn_backend=existing_backend) worker.topk = 1 worker.speculative_num_steps = 2 + worker.seed_dsa_topk_from_draft_extend = False with patch( "sglang.srt.speculative.eagle_worker_v2.DraftBackendFactory", @@ -151,10 +226,14 @@ class TestEagleWorkerV2BackendFallback(CustomTestCase): worker.draft_runner = SimpleNamespace(attn_backend=existing_backend) worker.topk = 1 worker.speculative_num_steps = 2 + worker.seed_dsa_topk_from_draft_extend = True + factory_kwargs = {} with patch( "sglang.srt.speculative.eagle_worker_v2.DraftBackendFactory", - _make_backend_factory(decode_backend, draft_extend_backend), + _make_backend_factory( + decode_backend, draft_extend_backend, captured_kwargs=factory_kwargs + ), ): worker.init_attention_backend() @@ -162,6 +241,7 @@ class TestEagleWorkerV2BackendFallback(CustomTestCase): self.assertIs(worker.draft_extend_attn_backend, draft_extend_backend) self.assertIs(worker.draft_runner.draft_attn_backend, decode_backend) self.assertIs(worker.draft_runner.attn_backend, draft_extend_backend) + self.assertTrue(factory_kwargs["seed_dsa_topk_from_draft_extend"]) def _make_adaptive_worker(self, runner_attn_backend): """An EAGLEWorkerV2 with a draft worker whose state-machine fields are