Use runtime token widths for Triton speculative verification (#39859)
Co-authored-by: raghotham <853234+raghotham@users.noreply.github.com>
This commit is contained in:
co-authored by
raghotham
parent
6bd1a0af1d
commit
248c202b46
@@ -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),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"]))
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user