[Spec] Rename token resolver to _resolve_spec_v2_tokens; remove dead V1 helpers (#27552)

This commit is contained in:
Liangsheng Yin
2026-06-08 14:42:21 -07:00
committed by GitHub
parent ca66e6fb5e
commit b5c64b94d5
6 changed files with 4 additions and 526 deletions
@@ -524,12 +524,12 @@ class SchedulerBatchResultProcessor:
logprob_pt += num_input_logprobs
return logprob_pt
def _resolve_spec_overlap_tokens(
def _resolve_spec_v2_tokens(
self,
result: GenerationBatchResult,
batch: ScheduleBatch,
) -> List[List[int]]:
"""Resolve the padding next token ids for speculative decoding with overlap."""
"""Resolve the padded next token ids for spec-v2 (overlap and non-overlap)."""
assert result.next_token_ids.is_cpu
assert result.accept_lens.is_cpu
@@ -712,7 +712,7 @@ class SchedulerBatchResultProcessor:
next_token_logprobs = None
if batch.spec_algorithm.is_none() or batch.is_spec_v2:
if batch.is_spec_v2:
next_token_ids = self._resolve_spec_overlap_tokens(result, batch)
next_token_ids = self._resolve_spec_v2_tokens(result, batch)
elif isinstance(next_token_ids, list):
pass # MLX path: already a list[int], skip torch round-trip
else:
@@ -43,7 +43,7 @@ logger = logging.getLogger(__name__)
def _get_draft_model_runner(draft_worker):
# DFlashWorker: exposes draft_model_runner directly
# DFlash / FrozenKVMTP workers expose draft_model_runner directly
runner = getattr(draft_worker, "draft_model_runner", None)
if runner is not None:
return runner
@@ -48,30 +48,6 @@ def per_step_draft_out_cache_loc(
)
def apply_eagle_prefill_input_rotation(
batch: ScheduleBatch, next_token_ids: torch.Tensor
) -> None:
"""EAGLE input rotation for draft prefill.
Each req's slice [t_0..t_{n-1}] -> [t_1..t_{n-1}, t_n] with
t_n = next_token_ids[i]. Aligns draft's position-i hidden with
target's label at i+1 — the basis of EAGLE chain prediction.
Vectorized: one whole-tensor left shift + scatter at segment tails.
"""
if batch.forward_mode.is_idle():
return
assert len(next_token_ids) == len(batch.seq_lens)
extend_lens = torch.tensor(
batch.extend_lens, dtype=torch.int64, device=batch.device
)
seg_ends = extend_lens.cumsum(0) - 1
rotated = torch.empty_like(batch.input_ids)
rotated[:-1] = batch.input_ids[1:]
# TODO: chunked-prefill chain divergence at non-final-chunk seg end; fix per PR #26329.
rotated[seg_ends] = next_token_ids.to(batch.input_ids.dtype)
batch.input_ids = rotated
def _eagle_prefill_tail_tokens(
batch: ScheduleBatch, next_token_ids: torch.Tensor
) -> torch.Tensor:
@@ -14,14 +14,10 @@ from sglang.srt.distributed.parallel_state import (
patch_tensor_parallel_group,
)
from sglang.srt.environ import envs
from sglang.srt.mem_cache.common import get_last_loc
from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.triton_ops.cache_locs import (
align_evict_mask_to_page_size as align_evict_mask_to_page_size,
)
from sglang.srt.speculative.triton_ops.cache_locs import (
assign_draft_cache_locs as assign_draft_cache_locs,
)
from sglang.srt.speculative.triton_ops.cache_locs import (
assign_req_to_token_pool as assign_req_to_token_pool,
)
@@ -468,39 +464,3 @@ def draft_tp_context(tp_group: GroupCoordinator):
# We disable mscclpp now because it doesn't support 2 comm groups.
with patch_tensor_parallel_group(tp_group):
yield
# Disable torch.compile for this function because it will be
# even slower.
# @torch.compile(dynamic=True)
def get_last_loc_large_page_size_large_top_k(
req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
speculative_num_steps: int,
topk: int,
page_size: int,
):
prefix_lens = seq_lens
last_page_lens = prefix_lens % page_size
num_new_pages_per_topk = (
last_page_lens + speculative_num_steps + page_size - 1
) // page_size
seq_lens = prefix_lens // page_size * page_size + num_new_pages_per_topk * (
page_size * topk
)
extend_lens = seq_lens - prefix_lens
last_loc = get_last_loc(
req_to_token,
req_pool_indices,
prefix_lens,
)
return (
prefix_lens,
seq_lens,
last_loc,
num_new_pages_per_topk,
extend_lens,
last_page_lens,
)
@@ -93,116 +93,6 @@ def assign_req_to_token_pool_func(
)
@triton.jit
def assign_draft_cache_locs(
req_pool_indices,
req_to_token,
seq_lens,
extend_lens,
num_new_pages_per_topk,
out_cache_loc,
source_cache_loc,
target_cache_loc,
last_page_lens_cumsum,
duplicate_cache_len: tl.constexpr,
pool_len: tl.constexpr,
topk: tl.constexpr,
speculative_num_steps: tl.constexpr,
page_size: tl.constexpr,
bs_upper: tl.constexpr,
iter_upper: tl.constexpr,
):
BLOCK_SIZE: tl.constexpr = 128
pid = tl.program_id(axis=0)
if page_size == 1 or topk == 1:
copy_len = topk * speculative_num_steps
out_cache_ptr = out_cache_loc + pid * topk * speculative_num_steps
else:
bs_offset = tl.arange(0, bs_upper)
copy_len = tl.load(extend_lens + pid)
cum_copy_len = tl.sum(tl.load(extend_lens + bs_offset, mask=bs_offset < pid))
out_cache_ptr = out_cache_loc + cum_copy_len
# Part 1: Copy from out_cache_loc to req_to_token
kv_start = tl.load(seq_lens + pid)
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
num_loop = tl.cdiv(copy_len, BLOCK_SIZE)
for i in range(num_loop):
copy_offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
mask = copy_offset < copy_len
data = tl.load(out_cache_ptr + copy_offset, mask=mask)
tl.store(token_pool + kv_start + copy_offset, data, mask=mask)
# XXX (MUSA): Triton issue: chained boolean operators (A or B or C) are not supported.
if (page_size != 1 and topk != 1) and duplicate_cache_len > 0:
# Part 2: Copy indices into source_cache_loc and target_cache_loc
# Expected output: src:[8,9,10,8,9,10...] tgt:[16,17,18,24,25,26...]
prefix_len = tl.load(seq_lens + pid)
last_page_len = prefix_len % page_size
offsets = tl.arange(0, page_size)
mask = offsets < last_page_len
num_new_pages_per_topk_ = tl.load(num_new_pages_per_topk + pid)
prefix_base = token_pool + prefix_len - last_page_len
src_indices = tl.load(prefix_base + offsets, mask=mask)
last_page_lens_cumsum_ = tl.load(last_page_lens_cumsum + pid)
# Skip the first one since no copy is needed
for topk_id in range(1, topk):
tl.store(
source_cache_loc
+ (topk - 1) * (last_page_lens_cumsum_ - last_page_len)
+ (topk_id - 1) * last_page_len
+ offsets,
src_indices,
mask=mask,
)
tgt_indices = tl.load(
prefix_base + topk_id * num_new_pages_per_topk_ * page_size + offsets,
mask=mask,
)
tl.store(
target_cache_loc
+ (topk - 1) * (last_page_lens_cumsum_ - last_page_len)
+ (topk_id - 1) * last_page_len
+ offsets,
tgt_indices,
mask=mask,
)
# Part 3: Copy and remove the used indices for duplication
# speculative_num_steps=5, page_size=4, num_new_pages_per_topk_=2, last_page_len=1
# - xxxxx .. | - xxxxx .. |
# topk=0 topk=1
# "-" means prefix tokens
# "x" means speculative draft tokens
# "." means padded tokens
# we only want to copy the "x" part.
iter_offset = tl.arange(0, iter_upper)
for topk_id in range(topk):
mask_upper = iter_offset < (speculative_num_steps + last_page_len)
mask_lower = iter_offset >= last_page_len
combined_mask = mask_upper & mask_lower
indices = tl.load(
prefix_base
+ topk_id * num_new_pages_per_topk_ * page_size
+ iter_offset,
mask=combined_mask,
other=0,
)
# Shift from previous batches
ptr_offset = pid * speculative_num_steps * topk
# Subtract last_page_len to fill the gap of duplicated last page tokens.
# For example, token pool is (1, 2, 3, 4 ,5) and last page is 1,
# we write 2, 3, 4 to the front of out_cache_loc.
tl.store(
out_cache_loc
+ ptr_offset
+ topk_id * speculative_num_steps
- last_page_len
+ iter_offset,
indices,
mask=combined_mask,
)
@triton.jit
def assign_draft_cache_locs_page_size_1(
req_pool_indices,