From 248c202b46d4a44ad46a3f09d48661b0d9ce6257 Mon Sep 17 00:00:00 2001 From: Lianmin Zheng Date: Fri, 18 Sep 2026 11:07:55 -0700 Subject: [PATCH] Use runtime token widths for Triton speculative verification (#39859) Co-authored-by: raghotham <853234+raghotham@users.noreply.github.com> --- .../srt/layers/attention/triton_backend.py | 96 +++---- python/sglang/srt/speculative/eagle_info.py | 6 +- python/sglang/srt/speculative/spec_info.py | 4 + .../attention_methods/gdn_attention.py | 2 + .../attention_methods/kda_attention.py | 2 + .../attention_methods/lightning_attention.py | 2 + .../attention_methods/mamba2_attention.py | 2 + .../attention_methods/mla_attention.py | 2 + .../attention/test_triton_verify_metadata.py | 237 ++++++++++++++++++ .../test_eagle_worker_v2_topk1_fastpath.py | 14 ++ 10 files changed, 302 insertions(+), 65 deletions(-) create mode 100644 test/registered/unit/layers/attention/test_triton_verify_metadata.py diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 239b57103..6cbfc87f0 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -213,6 +213,7 @@ class TritonAttnBackend(AttentionBackend): self.page_size = getattr(model_runner, "page_size", 1) or 1 self.kv_index_translator = model_runner.kv_index_translator self.num_draft_tokens = get_spec().speculative_num_draft_tokens + self.target_verify_num_tokens_per_req = model_runner.decode_num_tokens_per_req() self.speculative_num_steps = get_spec().speculative_num_steps self.topk = get_spec().speculative_eagle_topk or 0 # Split-KV verify is bit-equivalent only for a pure-causal chain (topk==1) @@ -551,6 +552,12 @@ class TritonAttnBackend(AttentionBackend): ) return kv_indptr, window_kv_indptr, window_kv_lens, num_kv_splits_lens + def _target_verify_num_tokens_per_req(self, spec_info: Optional[SpecInput]) -> int: + # Runtime metadata may vary by step; nonpositive means use capture width. + if spec_info is None or spec_info.num_tokens_per_req <= 0: + return self.target_verify_num_tokens_per_req + return spec_info.num_tokens_per_req + def _update_target_verify_buffers( self, bs: int, @@ -559,19 +566,12 @@ class TritonAttnBackend(AttentionBackend): req_pool_indices: torch.Tensor, ): """Fill all cuda-graph buffers for target_verify mode.""" - # Prefer the spec_info's per-request query length (DSpark draft propose - # uses gamma < verify window); fall back to the configured verify window. - num_draft_tokens = self.num_draft_tokens - if ( - spec_info is not None - and getattr(spec_info, "draft_token_num", None) is not None - ): - num_draft_tokens = int(spec_info.draft_token_num) + num_tokens_per_req = self._target_verify_num_tokens_per_req(spec_info) qo_indptr = self.qo_indptr[: bs + 1] qo_indptr[: bs + 1] = torch.arange( 0, - (1 + bs) * num_draft_tokens, - step=num_draft_tokens, + (1 + bs) * num_tokens_per_req, + step=num_tokens_per_req, dtype=torch.int32, device=self.device, ) @@ -601,14 +601,11 @@ class TritonAttnBackend(AttentionBackend): custom_mask = ( self._verify_mask.buffer if self._verify_mask is not None else None ) - if ( - spec_info is not None - and getattr(spec_info, "custom_mask", None) is not None - ): + if spec_info is not None and spec_info.custom_mask is not None: custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask else: custom_mask = None - seq_mask_len = num_draft_tokens * (seq_lens + num_draft_tokens) + seq_mask_len = num_tokens_per_req * (seq_lens + num_tokens_per_req) mask_indptr = self.mask_indptr[: bs + 1] mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0) return ( @@ -890,18 +887,11 @@ class TritonAttnBackend(AttentionBackend): max_extend_len = None elif forward_batch.forward_mode.is_target_verify(): bs = len(forward_batch.req_pool_indices) - # self.num_draft_tokens is the verify window (gamma + 1), while - # DSpark draft propose runs a gamma-token TARGET_VERIFY forward. - num_draft_tokens = self.num_draft_tokens - if ( - spec_info is not None - and getattr(spec_info, "draft_token_num", None) is not None - ): - num_draft_tokens = int(spec_info.draft_token_num) + num_tokens_per_req = self._target_verify_num_tokens_per_req(spec_info) qo_indptr = torch.arange( 0, - (1 + bs) * num_draft_tokens, - step=num_draft_tokens, + (1 + bs) * num_tokens_per_req, + step=num_tokens_per_req, dtype=torch.int32, device=self.device, ) @@ -938,13 +928,13 @@ class TritonAttnBackend(AttentionBackend): ) custom_mask = spec_info.custom_mask - seq_mask_len = num_draft_tokens * ( - forward_batch.seq_lens + num_draft_tokens + seq_mask_len = num_tokens_per_req * ( + forward_batch.seq_lens + num_tokens_per_req ) mask_indptr = self.mask_indptr mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0) mask_indptr = mask_indptr[: bs + 1] - max_extend_len = num_draft_tokens + max_extend_len = num_tokens_per_req num_kv_splits = None attn_logits = None attn_lse = None @@ -1208,15 +1198,10 @@ class TritonAttnBackend(AttentionBackend): self._verify_mask.buffer if self._verify_mask is not None and spec_info is not None - and getattr(spec_info, "custom_mask", None) is not None + and spec_info.custom_mask is not None else None ) - max_extend_len = self.num_draft_tokens - if ( - spec_info is not None - and getattr(spec_info, "draft_token_num", None) is not None - ): - max_extend_len = int(spec_info.draft_token_num) + max_extend_len = self._target_verify_num_tokens_per_req(spec_info) return ForwardMetadata( attn_logits=None, attn_lse=None, @@ -2008,23 +1993,14 @@ class TritonAttnBackend(AttentionBackend): # Compute window start positions (absolute position of first key in window) # window_start_pos = seq_len - window_len window_kv_lens = prefix_kv_indptr[1 : bs + 1] - prefix_kv_indptr[:bs] - # Handle TARGET_VERIFY mode where extend_prefix_lens might not be set if forward_batch.extend_prefix_lens is not None: window_start_pos = ( forward_batch.extend_prefix_lens[:bs] - window_kv_lens ) + elif forward_batch.forward_mode.is_target_verify(): + window_start_pos = forward_batch.seq_lens[:bs] - window_kv_lens else: - # Infer from spec_info: prefix_len = seq_len - draft_token_num - if forward_batch.spec_info is not None and hasattr( - forward_batch.spec_info, "draft_token_num" - ): - extend_prefix_lens = ( - forward_batch.seq_lens[:bs] - - forward_batch.spec_info.draft_token_num - ) - window_start_pos = extend_prefix_lens - window_kv_lens - else: - window_start_pos = None + window_start_pos = None else: sliding_window_size = -1 prefix_kv_indptr = self.forward_metadata.kv_indptr @@ -2047,29 +2023,23 @@ class TritonAttnBackend(AttentionBackend): elif self.forward_metadata.out_cache_loc_full_physical is not None: extend_kv_indices = self.forward_metadata.out_cache_loc_full_physical - # Handle cases where extend_seq_lens or extend_start_loc might not be set - # In speculative decoding, we can infer these from spec_info or compute them + # Capture batches may not have a spec_info, so use the attention + # metadata's resolved uniform verify width when extend lengths are absent. if forward_batch.extend_seq_lens is None: - # TARGET_VERIFY mode: infer extend_seq_lens from spec_info - if forward_batch.spec_info is not None and hasattr( - forward_batch.spec_info, "draft_token_num" - ): - draft_token_num = forward_batch.spec_info.draft_token_num - extend_seq_lens = torch.full( - (bs,), draft_token_num, dtype=torch.int32, device=self.device - ) - else: + if not forward_batch.forward_mode.is_target_verify(): raise RuntimeError( - "extend_seq_lens is None but cannot infer from spec_info. " - "This should not happen in TARGET_VERIFY mode." + "extend_seq_lens is None outside TARGET_VERIFY mode." ) + extend_seq_lens = torch.full( + (bs,), + self.forward_metadata.max_extend_len, + dtype=torch.int32, + device=self.device, + ) else: extend_seq_lens = forward_batch.extend_seq_lens - # Check extend_start_loc separately - it might be None even when extend_seq_lens is set if forward_batch.extend_start_loc is None: - # Compute extend_start_loc from extend_seq_lens - # extend_start_loc[i] = sum(extend_seq_lens[0:i]) extend_start_loc = torch.cat( [ torch.zeros(1, dtype=torch.int32, device=self.device), diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index e51f3f901..7c58117b3 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -1,5 +1,5 @@ import logging -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import List, Optional import torch @@ -15,7 +15,9 @@ logger = logging.getLogger(__name__) @dataclass class EagleVerifyInput(SpecInput): draft_token: torch.Tensor - custom_mask: torch.Tensor + # Keep this dataclass argument required despite SpecInput's None default; + # otherwise the required positions field would follow a defaulted field. + custom_mask: torch.Tensor = field() positions: torch.Tensor retrieve_index: torch.Tensor retrieve_next_token: torch.Tensor diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index 006d8244a..f296e1fe3 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -403,6 +403,10 @@ class SpecInput(ABC): num_tokens_per_req: int = -1 num_tokens_for_logprob_per_req: int = -1 + # Dataclasses assign fields before __post_init__ calls this base's __init__; + # assigning None there would overwrite the constructor's custom_mask. + custom_mask: Optional[torch.Tensor] = 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. diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py index 83b07097d..3595df14c 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py @@ -31,6 +31,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.runtime_context import get_context, get_parallel, get_server_args +from sglang.srt.speculative.spec_info import SpeculativeAlgorithm _parallel_override = get_parallel().override(attn_tp_size=1) _parallel_override.__enter__() @@ -227,6 +228,7 @@ class MockGDNModelRunner(ModelRunner): self.draft_attention_backend = None self.gpu_id = 0 self.ps = ParallelState.trivial() + self.spec_algorithm = SpeculativeAlgorithm.NONE self.canary_manager = None self.page_size = case.page_size self.model_config = model_config diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py index 00ddbc0a5..4dc11d16c 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py @@ -31,6 +31,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.runtime_context import get_context, get_parallel +from sglang.srt.speculative.spec_info import SpeculativeAlgorithm _parallel_override = get_parallel().override(attn_tp_size=1) _parallel_override.__enter__() @@ -230,6 +231,7 @@ class MockKDAModelRunner(ModelRunner): self.draft_attention_backend = None self.gpu_id = 0 self.ps = ParallelState.trivial() + self.spec_algorithm = SpeculativeAlgorithm.NONE self.canary_manager = None self.page_size = case.page_size self.model_config = model_config diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py index 8488d04bf..471750c7f 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py @@ -30,6 +30,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.runtime_context import get_context, get_parallel +from sglang.srt.speculative.spec_info import SpeculativeAlgorithm _parallel_override = get_parallel().override(attn_tp_size=1, attn_tp_rank=0) _parallel_override.__enter__() @@ -238,6 +239,7 @@ class MockLightningModelRunner(ModelRunner): self.draft_attention_backend = None self.gpu_id = 0 self.ps = ParallelState.trivial() + self.spec_algorithm = SpeculativeAlgorithm.NONE self.canary_manager = None self.page_size = case.page_size self.model_config = model_config diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py index 65b4ebe9b..b5fc99a51 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py @@ -51,6 +51,7 @@ from sglang.srt.model_executor.forward_context import ( # noqa: E402 forward_context, ) from sglang.srt.model_executor.model_runner import ModelRunner # noqa: E402 +from sglang.srt.speculative.spec_info import SpeculativeAlgorithm # noqa: E402 # Tiny dims chosen to be the minimum that satisfies MambaMixer2's TP/chunk asserts: # - num_heads % tp_size == 0 (tp_size=1) @@ -323,6 +324,7 @@ class MockMamba2ModelRunner(ModelRunner): self.draft_attention_backend = None self.gpu_id = 0 self.ps = ParallelState.trivial() + self.spec_algorithm = SpeculativeAlgorithm.NONE self.canary_manager = None self.page_size = case.page_size self.model_config = model_config diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py index 89cc7742f..8065fbb8a 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py @@ -25,6 +25,7 @@ from sglang.srt.model_executor.forward_context import ( from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.runtime_context import get_context, get_parallel +from sglang.srt.speculative.spec_info import SpeculativeAlgorithm _parallel_override = get_parallel().override(attn_tp_size=1) _parallel_override.__enter__() @@ -246,6 +247,7 @@ class MockMLAModelRunner(ModelRunner): self.dp_size = 1 self.pp_size = 1 self.ps = ParallelState.trivial() + self.spec_algorithm = SpeculativeAlgorithm.NONE speculative_num_draft_tokens = ( max(case.input_lens) if case.forward_mode.is_target_verify() diff --git a/test/registered/unit/layers/attention/test_triton_verify_metadata.py b/test/registered/unit/layers/attention/test_triton_verify_metadata.py new file mode 100644 index 000000000..fc3a46b86 --- /dev/null +++ b/test/registered/unit/layers/attention/test_triton_verify_metadata.py @@ -0,0 +1,237 @@ +import sys +from types import SimpleNamespace + +import pytest +import torch + +from sglang.srt.layers.attention.triton_backend import ( + ForwardMetadata, + TritonAttnBackend, +) +from sglang.srt.layers.radix_attention import AttentionType +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.speculative.spec_info import SpecInput, SpecInputType +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class _BonusTokenVerifyInput(SpecInput): + def __init__(self): + super().__init__(SpecInputType.EAGLE_VERIFY) + self.draft_token_num = 6 + self.num_tokens_per_req = 7 + + +class _KVIndexTranslator: + is_translating = False + + def fill_packed_read_stream( + self, + *, + req_pool_indices, + seq_lens, + indptr, + total_tokens, + out, + ): + out.zero_() + + +class _RecordingTritonBackend(TritonAttnBackend): + def build_unified_kv_indices( + self, + _prefix_kv_indptr, + _prefix_kv_indices, + extend_start_loc, + extend_seq_lens, + _extend_kv_indices, + batch_size, + ): + self.recorded_extend_start_loc = extend_start_loc.clone() + self.recorded_extend_seq_lens = extend_seq_lens.clone() + return ( + torch.zeros(batch_size + 1, dtype=torch.int32), + torch.zeros(batch_size * 7, dtype=torch.int64), + torch.zeros(batch_size, dtype=torch.int32), + ) + + def extend_attention_fwd_unified(self, *_args, **_kwargs): + self.recorded_window_start_pos = _kwargs["window_start_pos"].clone() + return None + + +def _make_backend(batch_size, capture_width=7): + backend = TritonAttnBackend.__new__(TritonAttnBackend) + backend.device = torch.device("cpu") + backend.num_draft_tokens = 9 + backend.target_verify_num_tokens_per_req = capture_width + backend.max_context_len = 128 + backend.qo_indptr = torch.zeros(batch_size + 1, dtype=torch.int32) + backend.kv_indptr = torch.zeros(batch_size + 1, dtype=torch.int32) + backend.window_kv_indptr = torch.zeros(batch_size + 1, dtype=torch.int32) + backend.mask_indptr = torch.zeros(batch_size + 1, dtype=torch.int64) + backend.cuda_graph_kv_indices = torch.zeros(256, dtype=torch.int64) + backend.kv_index_translator = _KVIndexTranslator() + backend.sliding_window_size = None + backend.use_sliding_window_kv_pool = False + backend._verify_mask = None + return backend + + +def _make_forward_batch(batch_size, spec_info): + seq_lens = torch.arange(32, 32 + batch_size, dtype=torch.int64) + return SimpleNamespace( + batch_size=batch_size, + input_ids=torch.zeros(batch_size * 7, dtype=torch.int64), + req_pool_indices=torch.arange(batch_size, dtype=torch.int64), + seq_lens=seq_lens, + seq_lens_cpu=seq_lens, + seq_lens_sum=int(seq_lens.sum()), + encoder_lens=None, + spec_info=spec_info, + forward_mode=ForwardMode.TARGET_VERIFY, + out_cache_loc=torch.zeros(batch_size * 7, dtype=torch.int64), + ) + + +def _assert_verify_width(backend, batch_size): + assert torch.equal( + backend.forward_metadata.qo_indptr, + torch.arange(0, (batch_size + 1) * 7, 7, dtype=torch.int32), + ) + assert backend.forward_metadata.max_extend_len == 7 + + +def test_eager_target_verify_uses_bonus_token_width(): + batch_size = 2 + backend = _make_backend(batch_size, capture_width=9) + spec_info = _BonusTokenVerifyInput() + + backend.init_forward_metadata(_make_forward_batch(batch_size, spec_info)) + + assert spec_info.draft_token_num == 6 + assert spec_info.num_tokens_per_req == 7 + _assert_verify_width(backend, batch_size) + assert torch.equal( + backend.forward_metadata.mask_indptr, + torch.tensor([0, 273, 553], dtype=torch.int64), + ) + + +@pytest.mark.parametrize( + ("batch_size", "raw_batch_size"), + ((2, 2), (4, 3)), +) +@pytest.mark.parametrize("unset_width", (None, -1, 0)) +def test_graph_capture_and_padded_replay_use_bonus_token_width( + batch_size, raw_batch_size, unset_width +): + backend = _make_backend(batch_size) + capture_spec = None + if unset_width is not None: + capture_spec = SpecInput(SpecInputType.EAGLE_VERIFY) + capture_spec.num_tokens_per_req = unset_width + capture_batch = _make_forward_batch(batch_size, capture_spec) + + backend.init_forward_metadata_out_graph( + capture_batch, + in_capture=True, + ) + _assert_verify_width(backend, batch_size) + + spec_info = _BonusTokenVerifyInput() + replay_batch = _make_forward_batch(batch_size, spec_info) + replay_batch.seq_lens[raw_batch_size:] = 1 + + backend.init_forward_metadata_out_graph( + replay_batch, + in_capture=False, + ) + + _assert_verify_width(backend, batch_size) + + +@pytest.mark.parametrize("with_spec_info", (False, True)) +def test_unified_target_verify_uses_resolved_metadata(with_spec_info): + batch_size = 2 + backend = _RecordingTritonBackend.__new__(_RecordingTritonBackend) + backend.device = torch.device("cpu") + backend.dcp_size = 1 + backend.enable_deterministic = True + backend.use_dense_fp8_chunked_prefill = False + backend.allow_bidirectional_attention_in_extend = False + backend.page_size = 1 + backend.token_to_kv_pool = SimpleNamespace( + get_key_buffer=lambda _layer_id: torch.zeros((1, 1, 4)), + get_value_buffer=lambda _layer_id: torch.zeros((1, 1, 4)), + ) + backend.forward_metadata = ForwardMetadata( + attn_logits=None, + attn_lse=None, + max_extend_len=7, + num_kv_splits=None, + kv_indptr=torch.zeros(batch_size + 1, dtype=torch.int32), + kv_indices=torch.zeros(1, dtype=torch.int64), + qo_indptr=torch.tensor([0, 7, 14], dtype=torch.int32), + custom_mask=None, + mask_indptr=None, + window_kv_indptr=torch.tensor([0, 8, 20], dtype=torch.int32), + window_kv_indices=torch.zeros(20, dtype=torch.int64), + window_num_kv_splits=None, + window_kv_offsets=None, + ) + layer = SimpleNamespace( + layer_id=0, + qk_head_dim=4, + v_head_dim=4, + tp_q_head_num=1, + k_scale=None, + v_scale=None, + logit_capping_method="tanh", + logit_cap=0.0, + is_cross_attention=False, + attn_type=AttentionType.DECODER, + sliding_window_size=16, + scaling=0.5, + xai_temperature_len=None, + ) + forward_batch = SimpleNamespace( + batch_size=batch_size, + forward_mode=ForwardMode.TARGET_VERIFY, + mha_one_shot=False, + out_cache_loc=torch.zeros(batch_size * 7, dtype=torch.int64), + extend_seq_lens=None, + extend_start_loc=None, + extend_prefix_lens=None, + seq_lens=torch.tensor([32, 64], dtype=torch.int64), + spec_info=_BonusTokenVerifyInput() if with_spec_info else None, + ) + q = torch.zeros((batch_size * 7, 4)) + + output = backend.forward_extend( + q, + q, + q, + layer, + forward_batch, + save_kv_cache=False, + ) + + assert output.shape == q.shape + assert torch.equal( + backend.recorded_extend_seq_lens, + torch.tensor([7, 7], dtype=torch.int32), + ) + assert torch.equal( + backend.recorded_extend_start_loc, + torch.tensor([0, 7], dtype=torch.int32), + ) + assert torch.equal( + backend.recorded_window_start_pos, + torch.tensor([24, 52], dtype=torch.int64), + ) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) 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 038fbf884..3f21212f1 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 @@ -16,6 +16,7 @@ import torch from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.runtime_context import get_context from sglang.srt.speculative.adaptive_runtime_state import SpecRuntimeState +from sglang.srt.speculative.eagle_info import EagleVerifyInput from sglang.srt.speculative.eagle_utils import organize_draft_results from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2 from sglang.test.ci.ci_register import register_amd_ci, register_cpu_ci @@ -135,6 +136,19 @@ class TestEagleWorkerV2Topk1FastPath(CustomTestCase): with self.assertRaises(AssertionError): worker._rebuild_topk1_chain_buffers() + def test_idle_verify_input_keeps_required_layout_tensors(self): + verify_input = EagleVerifyInput.create_idle_input( + topk=1, + spec_steps=3, + num_verify_tokens=4, + device=DEVICE, + ) + + self.assertEqual(verify_input.custom_mask.dtype, torch.bool) + self.assertEqual(verify_input.custom_mask.shape, (0,)) + self.assertEqual(verify_input.positions.dtype, torch.int64) + self.assertEqual(verify_input.positions.shape, (0,)) + def test_idle_draft_runs_each_eager_forward_without_tree_layout(self): worker = object.__new__(EagleDraftWorker) worker.speculative_num_steps = 3