[Spec] Rename token resolver to _resolve_spec_v2_tokens; remove dead V1 helpers (#27552)
This commit is contained in:
@@ -524,12 +524,12 @@ class SchedulerBatchResultProcessor:
|
|||||||
logprob_pt += num_input_logprobs
|
logprob_pt += num_input_logprobs
|
||||||
return logprob_pt
|
return logprob_pt
|
||||||
|
|
||||||
def _resolve_spec_overlap_tokens(
|
def _resolve_spec_v2_tokens(
|
||||||
self,
|
self,
|
||||||
result: GenerationBatchResult,
|
result: GenerationBatchResult,
|
||||||
batch: ScheduleBatch,
|
batch: ScheduleBatch,
|
||||||
) -> List[List[int]]:
|
) -> 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.next_token_ids.is_cpu
|
||||||
assert result.accept_lens.is_cpu
|
assert result.accept_lens.is_cpu
|
||||||
|
|
||||||
@@ -712,7 +712,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
next_token_logprobs = None
|
next_token_logprobs = None
|
||||||
if batch.spec_algorithm.is_none() or batch.is_spec_v2:
|
if batch.spec_algorithm.is_none() or batch.is_spec_v2:
|
||||||
if 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):
|
elif isinstance(next_token_ids, list):
|
||||||
pass # MLX path: already a list[int], skip torch round-trip
|
pass # MLX path: already a list[int], skip torch round-trip
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -43,7 +43,7 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
|
|
||||||
def _get_draft_model_runner(draft_worker):
|
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)
|
runner = getattr(draft_worker, "draft_model_runner", None)
|
||||||
if runner is not None:
|
if runner is not None:
|
||||||
return runner
|
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(
|
def _eagle_prefill_tail_tokens(
|
||||||
batch: ScheduleBatch, next_token_ids: torch.Tensor
|
batch: ScheduleBatch, next_token_ids: torch.Tensor
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|||||||
@@ -14,14 +14,10 @@ from sglang.srt.distributed.parallel_state import (
|
|||||||
patch_tensor_parallel_group,
|
patch_tensor_parallel_group,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
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.server_args import get_global_server_args
|
||||||
from sglang.srt.speculative.triton_ops.cache_locs import (
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
align_evict_mask_to_page_size as align_evict_mask_to_page_size,
|
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 (
|
from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||||
assign_req_to_token_pool as assign_req_to_token_pool,
|
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.
|
# We disable mscclpp now because it doesn't support 2 comm groups.
|
||||||
with patch_tensor_parallel_group(tp_group):
|
with patch_tensor_parallel_group(tp_group):
|
||||||
yield
|
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
|
@triton.jit
|
||||||
def assign_draft_cache_locs_page_size_1(
|
def assign_draft_cache_locs_page_size_1(
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
|
|||||||
@@ -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()
|
|
||||||
Reference in New Issue
Block a user