From b5c64b94d546490823fc91dac9b654450f962aeb Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Mon, 8 Jun 2026 14:42:21 -0700 Subject: [PATCH] [Spec] Rename token resolver to `_resolve_spec_v2_tokens`; remove dead V1 helpers (#27552) --- .../batch_result_processor.py | 6 +- .../scheduler_components/weight_updater.py | 2 +- python/sglang/srt/speculative/eagle_utils.py | 24 -- python/sglang/srt/speculative/spec_utils.py | 40 -- .../srt/speculative/triton_ops/cache_locs.py | 110 ------ test/manual/spec/test_spec_utils.py | 348 ------------------ 6 files changed, 4 insertions(+), 526 deletions(-) delete mode 100644 test/manual/spec/test_spec_utils.py diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index b0da5cd55..a28233958 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -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: diff --git a/python/sglang/srt/managers/scheduler_components/weight_updater.py b/python/sglang/srt/managers/scheduler_components/weight_updater.py index 17ab61793..77bf823b0 100644 --- a/python/sglang/srt/managers/scheduler_components/weight_updater.py +++ b/python/sglang/srt/managers/scheduler_components/weight_updater.py @@ -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 diff --git a/python/sglang/srt/speculative/eagle_utils.py b/python/sglang/srt/speculative/eagle_utils.py index 451a46dc9..3bb3cec3a 100644 --- a/python/sglang/srt/speculative/eagle_utils.py +++ b/python/sglang/srt/speculative/eagle_utils.py @@ -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: diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 80c742d66..abdee622b 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -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, - ) diff --git a/python/sglang/srt/speculative/triton_ops/cache_locs.py b/python/sglang/srt/speculative/triton_ops/cache_locs.py index 08edfb185..e8a6a754d 100644 --- a/python/sglang/srt/speculative/triton_ops/cache_locs.py +++ b/python/sglang/srt/speculative/triton_ops/cache_locs.py @@ -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, diff --git a/test/manual/spec/test_spec_utils.py b/test/manual/spec/test_spec_utils.py deleted file mode 100644 index 47085afde..000000000 --- a/test/manual/spec/test_spec_utils.py +++ /dev/null @@ -1,348 +0,0 @@ -import unittest - -import numpy as np -import torch - -from sglang.srt.mem_cache.memory_pool import copy_all_layer_kv_cache_tiled -from sglang.srt.speculative.spec_utils import assign_draft_cache_locs -from sglang.srt.utils import next_power_of_2 - -BYTES_PER_TILE = 128 - - -class TestSpecUtils(unittest.TestCase): - - def setUp(self): - self.device = "cuda" if torch.cuda.is_available() else "cpu" - self.data_ptrs = torch.zeros(2, 1, dtype=torch.uint64, device=self.device) - self.k_cache = [ - torch.zeros((100, 1, 1), dtype=torch.float32, device=self.device) - ] - self.v_cache = [ - torch.zeros((100, 1, 1), dtype=torch.float32, device=self.device) - ] - self.k_cache[0][:11, 0, 0] = torch.tensor( - [0.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0], - dtype=torch.float32, - device=self.device, - ) - self.v_cache[0][:11, 0, 0] = torch.tensor( - [-0.0, -0.1, -0.2, -0.3, -0.4, -0.5, -0.6, -0.7, -0.8, -0.9, -1.0], - dtype=torch.float32, - device=self.device, - ) - self.data_ptrs[0, 0] = self.k_cache[0].data_ptr() - self.data_ptrs[1, 0] = self.v_cache[0].data_ptr() - - self.data_strides = torch.tensor( - [ - np.prod(x.shape[1:]) * x.dtype.itemsize - for x in self.k_cache + self.v_cache - ], - device=self.device, - dtype=torch.int64, - ) - - def test_assign_draft_cache_locs_single_seq(self): - # Testing Setup: req_to_token starting from 4 - # 4,5,6,7,{8,9,10}, 8,9,10 is the last partial page, 3 tokens < page_size=4 - # next kv cache will be stored starting 11,12,13... - device = self.device - num_seqs = 1 - page_size = 4 - speculative_num_steps = 5 - topk = 8 - seq_lens_num = 7 - extend_lens_num = 61 # includes the duplicated last page - req_pool_indices = torch.arange(num_seqs, dtype=torch.int32, device=device) - req_to_token = torch.zeros((num_seqs, 100), dtype=torch.int32, device=device) - req_to_token[0, :seq_lens_num] = torch.tensor( - [4, 5, 6, 7, 8, 9, 10], device=device - ) - seq_lens = torch.tensor([seq_lens_num], dtype=torch.int32, device=device) - extend_lens = torch.tensor([extend_lens_num], dtype=torch.int32, device=device) - num_new_pages_per_topk = torch.tensor([2], dtype=torch.int32, device=device) - out_cache_loc = torch.arange(11, 11 + extend_lens_num, device=device) - last_page_lens = torch.tensor([3], dtype=torch.int32, device=device) - last_page_lens_cumsum = torch.cumsum(last_page_lens, dim=0) - duplicate_cache_len = last_page_lens.sum().item() * (topk - 1) - target_cache_loc = torch.zeros( - duplicate_cache_len, dtype=torch.int32, device=device - ) - source_cache_loc = torch.zeros( - duplicate_cache_len, dtype=torch.int32, device=device - ) - assign_draft_cache_locs[(num_seqs,)]( - 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, - req_to_token.shape[1], - topk, - speculative_num_steps, - page_size, - next_power_of_2(num_seqs), - next_power_of_2(speculative_num_steps + page_size), - ) - - out_cache_loc = out_cache_loc[: num_seqs * topk * speculative_num_steps] - expected_source_cache_loc = torch.tensor( - [8, 9, 10] * (topk - 1), device=device, dtype=torch.int32 - ) - assert torch.allclose(source_cache_loc, expected_source_cache_loc) - - copy_all_layer_kv_cache_tiled[(len(self.data_ptrs),)]( - self.data_ptrs, - self.data_strides, - target_cache_loc, - source_cache_loc, - len(target_cache_loc), - next_power_of_2(len(target_cache_loc)), - BYTES_PER_TILE, - ) - assert torch.allclose( - self.k_cache[0][16:19, 0, 0], - torch.tensor( - [0.8, 0.9, 1.0], - dtype=torch.float32, - device=device, - ), - ) - assert torch.allclose( - self.v_cache[0][16:19, 0, 0], - torch.tensor( - [-0.8, -0.9, -1.0], - dtype=torch.float32, - device=device, - ), - ) - - def test_assign_draft_cache_locs_multi_seq(self): - device = self.device - num_seqs = 3 - page_size = 4 - speculative_num_steps = 5 - topk = 8 - req_pool_indices = torch.arange(num_seqs, dtype=torch.int32, device=device) - req_to_token = torch.zeros((num_seqs, 100), dtype=torch.int32, device=device) - seq_lens = torch.tensor([8, 7, 5], dtype=torch.int32, device=device) - extend_lens = torch.tensor([64, 64, 64], dtype=torch.int32, device=device) - num_new_pages_per_topk = torch.tensor( - [2, 2, 2], dtype=torch.int32, device=device - ) - req_to_token = torch.zeros((num_seqs, 100), dtype=torch.int32, device=device) - req_to_token[0, :8] = torch.tensor([4, 5, 6, 7, 8, 9, 10, 11], device=device) - req_to_token[1, :7] = torch.tensor([4, 5, 6, 7, 8, 9, 10], device=device) - req_to_token[2, :5] = torch.tensor([4, 5, 6, 7, 8], device=device) - last_page_lens = torch.tensor([0, 3, 1], dtype=torch.int32, device=device) - last_page_lens_cumsum = torch.cumsum(last_page_lens, dim=0) - duplicate_cache_len = last_page_lens.sum().item() * (topk - 1) - out_cache_loc = torch.arange( - 12, 12 + torch.sum(extend_lens), dtype=torch.int32, device=device - ) - target_cache_loc = torch.zeros( - duplicate_cache_len, dtype=torch.int32, device=device - ) - source_cache_loc = torch.zeros( - duplicate_cache_len, dtype=torch.int32, device=device - ) - assign_draft_cache_locs[(num_seqs,)]( - 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, - req_to_token.shape[1], - topk, - speculative_num_steps, - page_size, - next_power_of_2(num_seqs), - next_power_of_2(speculative_num_steps + page_size), - ) - out_cache_loc = out_cache_loc[: num_seqs * topk * speculative_num_steps] - # fmt: off - expected_out_cache_loc = torch.tensor([ - 12, 13, 14, 15, 16, - 20, 21, 22, 23, 24, - 28, 29, 30, 31, 32, - 36, 37, 38, 39, 40, - 44, 45, 46, 47, 48, - 52, 53, 54, 55, 56, - 60, 61, 62, 63, 64, - 68, 69, 70, 71, 72, - 76, 77, 78, 79, 80, - 84, 85, 86, 87, 88, - 92, 93, 94, 95, 96, - 100, 101, 102, 103, 104, - 108, 109, 110, 111, 112, - 116, 117, 118, 119, 120, - 124, 125, 126, 127, 128, - 132, 133, 134, 135, 136, - 140, 141, 142, 143, 144, - 148, 149, 150, 151, 152, - 156, 157, 158, 159, 160, - 164, 165, 166, 167, 168, - 172, 173, 174, 175, 176, - 180, 181, 182, 183, 184, - 188, 189, 190, 191, 192, - 196, 197, 198, 199, 200 - ], device=device, dtype=torch.int32) - expected_source_cache_loc = torch.tensor([8, 9, 10] * 7 + [8] * 7, device=device, dtype=torch.int32) - expected_target_cache_loc = torch.tensor([ - 81, 82, 83, 89, 90, 91, 97, 98, 99, 105, 106, 107, 113, 114, - 115, 121, 122, 123, 129, 130, 131, 147, 155, 163, 171, 179, 187, 195 - ], device=device, dtype=torch.int32) - # fmt: on - assert torch.allclose(out_cache_loc, expected_out_cache_loc) - assert torch.allclose(source_cache_loc, expected_source_cache_loc) - assert torch.allclose(target_cache_loc, expected_target_cache_loc) - copy_all_layer_kv_cache_tiled[(len(self.data_ptrs),)]( - self.data_ptrs, - self.data_strides, - target_cache_loc, - source_cache_loc, - len(target_cache_loc), - next_power_of_2(len(target_cache_loc)), - BYTES_PER_TILE, - ) - assert torch.allclose( - self.k_cache[0][81:84, 0, 0], - torch.tensor( - [0.8, 0.9, 1.0], - dtype=torch.float32, - device=device, - ), - ) - assert torch.allclose( - self.v_cache[0][81:84, 0, 0], - torch.tensor( - [-0.8, -0.9, -1.0], - dtype=torch.float32, - device=device, - ), - ) - - def test_assign_draft_cache_locs_page_size_1(self): - # Test to make sure page_size=1 not affected - device = self.device - num_seqs = 1 - page_size = 1 - speculative_num_steps = 5 - topk = 8 - seq_lens_num = 7 - extend_lens_num = topk * speculative_num_steps - req_pool_indices = torch.arange(num_seqs, dtype=torch.int32, device=device) - req_to_token = torch.zeros((num_seqs, 100), dtype=torch.int32, device=device) - req_to_token[0, :seq_lens_num] = torch.tensor( - [4, 5, 6, 7, 8, 9, 10], device=device - ) - seq_lens = torch.tensor([seq_lens_num], dtype=torch.int32, device=device) - extend_lens = torch.tensor([extend_lens_num], dtype=torch.int32, device=device) - num_new_pages_per_topk = torch.tensor([2], dtype=torch.int32, device=device) - out_cache_loc = torch.arange(11, 11 + extend_lens_num, device=device) - last_page_lens = torch.tensor([3], dtype=torch.int32, device=device) - duplicate_cache_len = 0 - target_cache_loc = None - source_cache_loc = None - last_page_lens_cumsum = None - assign_draft_cache_locs[(num_seqs,)]( - 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, - req_to_token.shape[1], - topk, - speculative_num_steps, - page_size, - next_power_of_2(num_seqs), - next_power_of_2(speculative_num_steps + page_size), - ) - out_cache_loc = out_cache_loc[: num_seqs * topk * speculative_num_steps] - expected_out_cache_loc = torch.arange(11, 11 + extend_lens_num, device=device) - assert torch.allclose(out_cache_loc, expected_out_cache_loc) - - def test_assign_draft_cache_locs_page_size_gt_spec_steps(self): - device = self.device - num_seqs = 1 - page_size = 16 - speculative_num_steps = 4 - topk = 3 - seq_lens_num = 12 - pool_len = 256 - req_pool_indices = torch.arange(num_seqs, dtype=torch.int32, device=device) - req_to_token = torch.zeros( - (num_seqs, pool_len), dtype=torch.int32, device=device - ) - req_to_token[0, :seq_lens_num] = torch.arange( - seq_lens_num, dtype=torch.int32, device=device - ) - seq_lens = torch.tensor([seq_lens_num], dtype=torch.int32, device=device) - last_page_len = seq_lens_num % page_size - last_page_lens = torch.tensor([last_page_len], dtype=torch.int32, device=device) - last_page_lens_cumsum = torch.cumsum(last_page_lens, dim=0) - num_new_pages_per_topk_val = ( - last_page_len + speculative_num_steps + page_size - 1 - ) // page_size - num_new_pages_per_topk = torch.tensor( - [num_new_pages_per_topk_val], dtype=torch.int32, device=device - ) - extend_lens_num = num_new_pages_per_topk_val * page_size * topk - extend_lens = torch.tensor([extend_lens_num], dtype=torch.int32, device=device) - out_cache_loc = torch.arange( - 2000, 2000 + extend_lens_num, dtype=torch.int32, device=device - ) - duplicate_cache_len = last_page_lens.sum().item() * (topk - 1) - target_cache_loc = torch.zeros( - duplicate_cache_len, dtype=torch.int32, device=device - ) - source_cache_loc = torch.zeros( - duplicate_cache_len, dtype=torch.int32, device=device - ) - assign_draft_cache_locs[(num_seqs,)]( - 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, - req_to_token.shape[1], - topk, - speculative_num_steps, - page_size, - next_power_of_2(num_seqs), - next_power_of_2(speculative_num_steps + page_size), - ) - trimmed = out_cache_loc[: num_seqs * topk * speculative_num_steps] - expected = [] - for topk_id in range(topk): - start = seq_lens_num + topk_id * num_new_pages_per_topk_val * page_size - expected.append( - req_to_token[0, start : start + speculative_num_steps].clone() - ) - expected_out_cache_loc = torch.cat(expected) - assert torch.allclose(trimmed, expected_out_cache_loc) - - -if __name__ == "__main__": - unittest.main()