Fix bookkeeping fields not encapsulated with real allocations in normal alloc, PD pre-alloc, DFlash and EAGLE (#29432)
This commit is contained in:
@@ -10,7 +10,6 @@ from sglang.kernels.spec import KernelBackend, KernelSpec
|
|||||||
|
|
||||||
# (module, public_fn) migrated from speculative/triton_ops.
|
# (module, public_fn) migrated from speculative/triton_ops.
|
||||||
_TRITON_KERNELS = [
|
_TRITON_KERNELS = [
|
||||||
("cache_locs", "assign_req_to_token_pool_func"),
|
|
||||||
("cache_locs", "assign_extend_cache_locs_func"),
|
("cache_locs", "assign_extend_cache_locs_func"),
|
||||||
("cache_locs", "generate_draft_decode_kv_indices"),
|
("cache_locs", "generate_draft_decode_kv_indices"),
|
||||||
("eagle", "fill_bonus_tokens"),
|
("eagle", "fill_bonus_tokens"),
|
||||||
|
|||||||
@@ -22,71 +22,7 @@ _is_musa = is_musa()
|
|||||||
_is_xpu = is_xpu()
|
_is_xpu = is_xpu()
|
||||||
|
|
||||||
if _is_cpu:
|
if _is_cpu:
|
||||||
from sgl_kernel import assign_extend_cache_locs_cpu, assign_req_to_token_pool_cpu
|
from sgl_kernel import assign_extend_cache_locs_cpu
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def assign_req_to_token_pool(
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
start_offset,
|
|
||||||
end_offset,
|
|
||||||
out_cache_loc,
|
|
||||||
pool_len: tl.constexpr,
|
|
||||||
bs_upper: tl.constexpr,
|
|
||||||
):
|
|
||||||
BLOCK_SIZE: tl.constexpr = 32
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
kv_start = tl.load(start_offset + pid)
|
|
||||||
kv_end = tl.load(end_offset + pid)
|
|
||||||
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
|
|
||||||
|
|
||||||
length_offset = tl.arange(0, bs_upper)
|
|
||||||
start = tl.load(start_offset + length_offset, mask=length_offset < pid, other=0)
|
|
||||||
end = tl.load(end_offset + length_offset, mask=length_offset < pid, other=0)
|
|
||||||
out_offset = tl.sum(end - start, axis=0)
|
|
||||||
|
|
||||||
out_cache_ptr = out_cache_loc + out_offset
|
|
||||||
|
|
||||||
save_offset = tl.arange(0, BLOCK_SIZE) + kv_start
|
|
||||||
load_offset = tl.arange(0, BLOCK_SIZE)
|
|
||||||
|
|
||||||
num_loop = tl.cdiv(kv_end - kv_start, BLOCK_SIZE)
|
|
||||||
for _ in range(num_loop):
|
|
||||||
mask = save_offset < kv_end
|
|
||||||
data = tl.load(out_cache_ptr + load_offset, mask=mask)
|
|
||||||
tl.store(token_pool + save_offset, data, mask=mask)
|
|
||||||
save_offset += BLOCK_SIZE
|
|
||||||
load_offset += BLOCK_SIZE
|
|
||||||
|
|
||||||
|
|
||||||
def assign_req_to_token_pool_func(
|
|
||||||
req_pool_indices: torch.Tensor,
|
|
||||||
req_to_token: torch.Tensor,
|
|
||||||
start_offset: torch.Tensor,
|
|
||||||
end_offset: torch.Tensor,
|
|
||||||
out_cache_loc: torch.Tensor,
|
|
||||||
batch_size: int,
|
|
||||||
):
|
|
||||||
if _is_cpu:
|
|
||||||
assign_req_to_token_pool_cpu(
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
start_offset,
|
|
||||||
end_offset,
|
|
||||||
out_cache_loc,
|
|
||||||
req_to_token.shape[1],
|
|
||||||
)
|
|
||||||
return
|
|
||||||
assign_req_to_token_pool[(batch_size,)](
|
|
||||||
req_pool_indices,
|
|
||||||
req_to_token,
|
|
||||||
start_offset,
|
|
||||||
end_offset,
|
|
||||||
out_cache_loc,
|
|
||||||
req_to_token.shape[1],
|
|
||||||
next_power_of_2(batch_size),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
|
|||||||
@@ -63,6 +63,7 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.managers.schedule_batch import (
|
from sglang.srt.managers.schedule_batch import (
|
||||||
FINISH_ABORT,
|
FINISH_ABORT,
|
||||||
NextBatchPlan,
|
NextBatchPlan,
|
||||||
|
ReqKvInfo,
|
||||||
ScheduleBatch,
|
ScheduleBatch,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.schedule_policy import match_prefix_for_req
|
from sglang.srt.managers.schedule_policy import match_prefix_for_req
|
||||||
@@ -1420,15 +1421,6 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
), "req_pool_indices is full! There is a bug in memory estimation."
|
), "req_pool_indices is full! There is a bug in memory estimation."
|
||||||
|
|
||||||
fill_len = self._pre_alloc_fill_len(req)
|
fill_len = self._pre_alloc_fill_len(req)
|
||||||
# TODO(th4): co-locate this req.kv bookkeeping with the real KV
|
|
||||||
# allocation; the pool alloc above and the kv_allocated_len assignment
|
|
||||||
# below should become a single owned-kv allocation step.
|
|
||||||
if req.kv is None:
|
|
||||||
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
|
||||||
|
|
||||||
req.kv = ReqKvInfo(kv_allocated_len=fill_len, swa_evicted_seqlen=0)
|
|
||||||
else:
|
|
||||||
req.kv.kv_allocated_len = fill_len
|
|
||||||
req.kv_committed_len = fill_len
|
req.kv_committed_len = fill_len
|
||||||
|
|
||||||
if prefix_len > 0:
|
if prefix_len > 0:
|
||||||
@@ -1466,6 +1458,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
f"req={req.rid}"
|
f"req={req.rid}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
allocator = self.token_to_kv_pool_allocator
|
||||||
if self.scheduler.enable_hisparse:
|
if self.scheduler.enable_hisparse:
|
||||||
# HiSparse is incompatible with decode-side L1 radix cache. Keep
|
# HiSparse is incompatible with decode-side L1 radix cache. Keep
|
||||||
# this path on the upstream full-allocation semantics.
|
# this path on the upstream full-allocation semantics.
|
||||||
@@ -1474,34 +1467,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
# Direct-to-host path: only allocate logical indices (no hisparse
|
# Direct-to-host path: only allocate logical indices (no hisparse
|
||||||
# device indices) and allocate host indices for RDMA destination.
|
# device indices) and allocate host indices for RDMA destination.
|
||||||
coordinator = self.scheduler.hisparse_coordinator
|
coordinator = self.scheduler.hisparse_coordinator
|
||||||
device = self.token_to_kv_pool_allocator.device
|
kv_loc = alloc_for_decode_prealloc_hisparse(
|
||||||
prefix_lens = torch.tensor([0], dtype=torch.int64, device=device)
|
allocator,
|
||||||
prefix_lens_cpu = torch.tensor([0], dtype=torch.int64)
|
req=req,
|
||||||
seq_lens = torch.tensor([fill_len], dtype=torch.int64, device=device)
|
fill_len=fill_len,
|
||||||
seq_lens_cpu = torch.tensor([fill_len], dtype=torch.int64)
|
uses_swa_tail=self._uses_swa_tail_prealloc(),
|
||||||
last_loc = torch.tensor([-1], dtype=torch.int64, device=device)
|
swa_tail_len=self._swa_tail_len(fill_len),
|
||||||
if self._uses_swa_tail_prealloc():
|
|
||||||
swa_tail_len = self._swa_tail_len(fill_len)
|
|
||||||
kv_loc = self.token_to_kv_pool_allocator.alloc_extend_swa_tail(
|
|
||||||
prefix_lens=prefix_lens,
|
|
||||||
prefix_lens_cpu=prefix_lens_cpu,
|
|
||||||
seq_lens=seq_lens,
|
|
||||||
seq_lens_cpu=seq_lens_cpu,
|
|
||||||
last_loc=last_loc,
|
|
||||||
extend_num_tokens=fill_len,
|
|
||||||
swa_tail_len=swa_tail_len,
|
|
||||||
)
|
)
|
||||||
req.swa_evicted_seqlen = fill_len - swa_tail_len
|
|
||||||
else:
|
|
||||||
kv_loc = self.token_to_kv_pool_allocator.alloc_logical_only(
|
|
||||||
prefix_lens=prefix_lens,
|
|
||||||
prefix_lens_cpu=prefix_lens_cpu,
|
|
||||||
seq_lens=seq_lens,
|
|
||||||
seq_lens_cpu=seq_lens_cpu,
|
|
||||||
last_loc=last_loc,
|
|
||||||
extend_num_tokens=fill_len,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Allocate host indices for the RDMA transfer target.
|
# Allocate host indices for the RDMA transfer target.
|
||||||
host_indices = coordinator.mem_pool_host.alloc_paged_token_slots(
|
host_indices = coordinator.mem_pool_host.alloc_paged_token_slots(
|
||||||
coordinator.req_to_host_pool,
|
coordinator.req_to_host_pool,
|
||||||
@@ -1510,42 +1482,20 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
0,
|
0,
|
||||||
coordinator.host_token_len(fill_len),
|
coordinator.host_token_len(fill_len),
|
||||||
)
|
)
|
||||||
elif self.token_to_kv_pool_allocator.page_size == 1:
|
|
||||||
kv_loc = self.token_to_kv_pool_allocator.alloc(delta_len)
|
|
||||||
else:
|
else:
|
||||||
device = self.token_to_kv_pool_allocator.device
|
uses_swa_tail = self._uses_swa_tail_prealloc() and prefix_len == 0
|
||||||
last_loc = (
|
swa_tail_len = self._swa_tail_len(fill_len)
|
||||||
prefix_indices[-1:].to(dtype=torch.int64, device=device)
|
kv_loc = alloc_for_decode_prealloc(
|
||||||
if prefix_len > 0
|
allocator,
|
||||||
else torch.tensor([-1], dtype=torch.int64, device=device)
|
req=req,
|
||||||
|
fill_len=fill_len,
|
||||||
|
delta_len=delta_len,
|
||||||
|
prefix_len=prefix_len,
|
||||||
|
total_prefix_len=total_prefix_len,
|
||||||
|
prefix_indices=prefix_indices,
|
||||||
|
uses_swa_tail=uses_swa_tail,
|
||||||
|
swa_tail_len=swa_tail_len,
|
||||||
)
|
)
|
||||||
if self._uses_swa_tail_prealloc() and prefix_len == 0:
|
|
||||||
# Tail-only SWA allocation: only valid when prefix_len == 0.
|
|
||||||
# When prefix_len > 0 (radix cache hit), we fall back to
|
|
||||||
# alloc_extend which allocates SWA at full page count; the
|
|
||||||
# SWA budget in that case may slightly under-estimate.
|
|
||||||
kv_loc = self.token_to_kv_pool_allocator.alloc_extend_swa_tail(
|
|
||||||
prefix_lens=torch.tensor([0], dtype=torch.int64, device=device),
|
|
||||||
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
|
||||||
seq_lens=torch.tensor([fill_len], dtype=torch.int64, device=device),
|
|
||||||
seq_lens_cpu=torch.tensor([fill_len], dtype=torch.int64),
|
|
||||||
last_loc=last_loc,
|
|
||||||
extend_num_tokens=fill_len,
|
|
||||||
swa_tail_len=self._swa_tail_len(fill_len),
|
|
||||||
)
|
|
||||||
req.kv.swa_evicted_seqlen = fill_len - self._swa_tail_len(fill_len)
|
|
||||||
else:
|
|
||||||
kv_loc = self.token_to_kv_pool_allocator.alloc_extend(
|
|
||||||
prefix_lens=torch.tensor(
|
|
||||||
[total_prefix_len], dtype=torch.int64, device=device
|
|
||||||
),
|
|
||||||
prefix_lens_cpu=torch.tensor([total_prefix_len], dtype=torch.int64),
|
|
||||||
seq_lens=torch.tensor([fill_len], dtype=torch.int64, device=device),
|
|
||||||
seq_lens_cpu=torch.tensor([fill_len], dtype=torch.int64),
|
|
||||||
last_loc=last_loc,
|
|
||||||
extend_num_tokens=delta_len,
|
|
||||||
)
|
|
||||||
|
|
||||||
assert kv_loc is not None, (
|
assert kv_loc is not None, (
|
||||||
f"KV cache is full! Bug in memory estimation. "
|
f"KV cache is full! Bug in memory estimation. "
|
||||||
f"available={self.token_to_kv_pool_allocator.available_size()}, "
|
f"available={self.token_to_kv_pool_allocator.available_size()}, "
|
||||||
@@ -1584,6 +1534,101 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
return kv_loc
|
return kv_loc
|
||||||
|
|
||||||
|
|
||||||
|
def alloc_for_decode_prealloc_hisparse(
|
||||||
|
allocator: BaseTokenToKVPoolAllocator,
|
||||||
|
*,
|
||||||
|
req: Req,
|
||||||
|
fill_len: int,
|
||||||
|
uses_swa_tail: bool,
|
||||||
|
swa_tail_len: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if req.kv is None:
|
||||||
|
req.kv = ReqKvInfo(kv_allocated_len=fill_len, swa_evicted_seqlen=0)
|
||||||
|
else:
|
||||||
|
req.kv.kv_allocated_len = fill_len
|
||||||
|
device = allocator.device
|
||||||
|
prefix_lens = torch.tensor([0], dtype=torch.int64, device=device)
|
||||||
|
prefix_lens_cpu = torch.tensor([0], dtype=torch.int64)
|
||||||
|
seq_lens = torch.tensor([fill_len], dtype=torch.int64, device=device)
|
||||||
|
seq_lens_cpu = torch.tensor([fill_len], dtype=torch.int64)
|
||||||
|
last_loc = torch.tensor([-1], dtype=torch.int64, device=device)
|
||||||
|
if uses_swa_tail:
|
||||||
|
kv_loc = allocator.alloc_extend_swa_tail(
|
||||||
|
prefix_lens=prefix_lens,
|
||||||
|
prefix_lens_cpu=prefix_lens_cpu,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
seq_lens_cpu=seq_lens_cpu,
|
||||||
|
last_loc=last_loc,
|
||||||
|
extend_num_tokens=fill_len,
|
||||||
|
swa_tail_len=swa_tail_len,
|
||||||
|
)
|
||||||
|
req.kv.swa_evicted_seqlen = fill_len - swa_tail_len
|
||||||
|
else:
|
||||||
|
kv_loc = allocator.alloc_logical_only(
|
||||||
|
prefix_lens=prefix_lens,
|
||||||
|
prefix_lens_cpu=prefix_lens_cpu,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
seq_lens_cpu=seq_lens_cpu,
|
||||||
|
last_loc=last_loc,
|
||||||
|
extend_num_tokens=fill_len,
|
||||||
|
)
|
||||||
|
return kv_loc
|
||||||
|
|
||||||
|
|
||||||
|
def alloc_for_decode_prealloc(
|
||||||
|
allocator: BaseTokenToKVPoolAllocator,
|
||||||
|
*,
|
||||||
|
req: Req,
|
||||||
|
fill_len: int,
|
||||||
|
delta_len: int,
|
||||||
|
prefix_len: int,
|
||||||
|
total_prefix_len: int,
|
||||||
|
prefix_indices: Optional[torch.Tensor],
|
||||||
|
uses_swa_tail: bool,
|
||||||
|
swa_tail_len: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if req.kv is None:
|
||||||
|
req.kv = ReqKvInfo(kv_allocated_len=fill_len, swa_evicted_seqlen=0)
|
||||||
|
else:
|
||||||
|
req.kv.kv_allocated_len = fill_len
|
||||||
|
if allocator.page_size == 1:
|
||||||
|
kv_loc = allocator.alloc(delta_len)
|
||||||
|
else:
|
||||||
|
device = allocator.device
|
||||||
|
last_loc = (
|
||||||
|
prefix_indices[-1:].to(dtype=torch.int64, device=device)
|
||||||
|
if prefix_len > 0
|
||||||
|
else torch.tensor([-1], dtype=torch.int64, device=device)
|
||||||
|
)
|
||||||
|
if uses_swa_tail:
|
||||||
|
# Tail-only SWA allocation: only valid when prefix_len == 0.
|
||||||
|
# When prefix_len > 0 (radix cache hit), we fall back to
|
||||||
|
# alloc_extend which allocates SWA at full page count; the
|
||||||
|
# SWA budget in that case may slightly under-estimate.
|
||||||
|
kv_loc = allocator.alloc_extend_swa_tail(
|
||||||
|
prefix_lens=torch.tensor([0], dtype=torch.int64, device=device),
|
||||||
|
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
||||||
|
seq_lens=torch.tensor([fill_len], dtype=torch.int64, device=device),
|
||||||
|
seq_lens_cpu=torch.tensor([fill_len], dtype=torch.int64),
|
||||||
|
last_loc=last_loc,
|
||||||
|
extend_num_tokens=fill_len,
|
||||||
|
swa_tail_len=swa_tail_len,
|
||||||
|
)
|
||||||
|
req.kv.swa_evicted_seqlen = fill_len - swa_tail_len
|
||||||
|
else:
|
||||||
|
kv_loc = allocator.alloc_extend(
|
||||||
|
prefix_lens=torch.tensor(
|
||||||
|
[total_prefix_len], dtype=torch.int64, device=device
|
||||||
|
),
|
||||||
|
prefix_lens_cpu=torch.tensor([total_prefix_len], dtype=torch.int64),
|
||||||
|
seq_lens=torch.tensor([fill_len], dtype=torch.int64, device=device),
|
||||||
|
seq_lens_cpu=torch.tensor([fill_len], dtype=torch.int64),
|
||||||
|
last_loc=last_loc,
|
||||||
|
extend_num_tokens=delta_len,
|
||||||
|
)
|
||||||
|
return kv_loc
|
||||||
|
|
||||||
|
|
||||||
class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
||||||
"""
|
"""
|
||||||
Store the requests that is polling kv
|
Store the requests that is polling kv
|
||||||
|
|||||||
@@ -2175,10 +2175,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# allocation in alloc_for_extend above; they are currently a few
|
# allocation in alloc_for_extend above; they are currently a few
|
||||||
# steps apart and should become one owned-kv allocation step.
|
# steps apart and should become one owned-kv allocation step.
|
||||||
req.kv_committed_len = seq_len
|
req.kv_committed_len = seq_len
|
||||||
if req.kv is None:
|
|
||||||
req.kv = ReqKvInfo(kv_allocated_len=seq_len, swa_evicted_seqlen=0)
|
|
||||||
else:
|
|
||||||
req.kv.kv_allocated_len = seq_len
|
|
||||||
|
|
||||||
# If input_embeds are available, store them
|
# If input_embeds are available, store them
|
||||||
if req.input_embeds is not None:
|
if req.input_embeds is not None:
|
||||||
@@ -2778,7 +2774,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
for req in self.reqs:
|
for req in self.reqs:
|
||||||
req.decode_batch_idx += 1
|
req.decode_batch_idx += 1
|
||||||
req.kv_committed_len += 1
|
req.kv_committed_len += 1
|
||||||
req.kv.kv_allocated_len += 1
|
|
||||||
|
|
||||||
# New-tensor avoids racing model_worker_batch refs queued for
|
# New-tensor avoids racing model_worker_batch refs queued for
|
||||||
# overlap forward.
|
# overlap forward.
|
||||||
|
|||||||
@@ -1,9 +1,12 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
from collections import defaultdict
|
||||||
from typing import TYPE_CHECKING, Optional
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.kernels.ops.memory.common import (
|
from sglang.kernels.ops.memory.common import (
|
||||||
get_last_loc_triton,
|
get_last_loc_triton,
|
||||||
@@ -24,12 +27,23 @@ from sglang.srt.mem_cache.common import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_server_args
|
||||||
from sglang.srt.utils import is_cuda, is_hip, is_npu, support_triton
|
from sglang.srt.utils import (
|
||||||
|
is_cpu,
|
||||||
|
is_cuda,
|
||||||
|
is_hip,
|
||||||
|
is_npu,
|
||||||
|
next_power_of_2,
|
||||||
|
support_triton,
|
||||||
|
)
|
||||||
from sglang.srt.utils.common import is_pin_memory_available
|
from sglang.srt.utils.common import is_pin_memory_available
|
||||||
|
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
|
_is_cpu = is_cpu()
|
||||||
|
|
||||||
|
if _is_cpu:
|
||||||
|
from sgl_kernel import assign_req_to_token_pool_cpu
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||||
@@ -360,6 +374,14 @@ def alloc_for_extend(
|
|||||||
batch.seq_lens_cpu,
|
batch.seq_lens_cpu,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||||
|
|
||||||
|
for req, seq_len in zip(batch.reqs, batch.seq_lens_cpu.tolist()):
|
||||||
|
if req.kv is None:
|
||||||
|
req.kv = ReqKvInfo(kv_allocated_len=seq_len, swa_evicted_seqlen=0)
|
||||||
|
else:
|
||||||
|
req.kv.kv_allocated_len = seq_len
|
||||||
|
|
||||||
return out_cache_loc, req_pool_indices_device, req_pool_indices_cpu
|
return out_cache_loc, req_pool_indices_device, req_pool_indices_cpu
|
||||||
|
|
||||||
|
|
||||||
@@ -466,4 +488,134 @@ def alloc_for_decode(batch: ScheduleBatch, token_per_req: int) -> torch.Tensor:
|
|||||||
token_per_req,
|
token_per_req,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
for req in batch.reqs:
|
||||||
|
req.kv.kv_allocated_len += token_per_req
|
||||||
|
|
||||||
return out_cache_loc
|
return out_cache_loc
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def assign_req_to_token_pool(
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
start_offset,
|
||||||
|
end_offset,
|
||||||
|
out_cache_loc,
|
||||||
|
pool_len: tl.constexpr,
|
||||||
|
bs_upper: tl.constexpr,
|
||||||
|
):
|
||||||
|
BLOCK_SIZE: tl.constexpr = 32
|
||||||
|
pid = tl.program_id(axis=0)
|
||||||
|
kv_start = tl.load(start_offset + pid)
|
||||||
|
kv_end = tl.load(end_offset + pid)
|
||||||
|
token_pool = req_to_token + tl.load(req_pool_indices + pid) * pool_len
|
||||||
|
|
||||||
|
length_offset = tl.arange(0, bs_upper)
|
||||||
|
start = tl.load(start_offset + length_offset, mask=length_offset < pid, other=0)
|
||||||
|
end = tl.load(end_offset + length_offset, mask=length_offset < pid, other=0)
|
||||||
|
out_offset = tl.sum(end - start, axis=0)
|
||||||
|
|
||||||
|
out_cache_ptr = out_cache_loc + out_offset
|
||||||
|
|
||||||
|
save_offset = tl.arange(0, BLOCK_SIZE) + kv_start
|
||||||
|
load_offset = tl.arange(0, BLOCK_SIZE)
|
||||||
|
|
||||||
|
num_loop = tl.cdiv(kv_end - kv_start, BLOCK_SIZE)
|
||||||
|
for _ in range(num_loop):
|
||||||
|
mask = save_offset < kv_end
|
||||||
|
data = tl.load(out_cache_ptr + load_offset, mask=mask)
|
||||||
|
tl.store(token_pool + save_offset, data, mask=mask)
|
||||||
|
save_offset += BLOCK_SIZE
|
||||||
|
load_offset += BLOCK_SIZE
|
||||||
|
|
||||||
|
|
||||||
|
def assign_req_to_token_pool_func(
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
start_offset: torch.Tensor,
|
||||||
|
end_offset: torch.Tensor,
|
||||||
|
out_cache_loc: torch.Tensor,
|
||||||
|
batch_size: int,
|
||||||
|
):
|
||||||
|
if _is_cpu:
|
||||||
|
assign_req_to_token_pool_cpu(
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
start_offset,
|
||||||
|
end_offset,
|
||||||
|
out_cache_loc,
|
||||||
|
req_to_token.shape[1],
|
||||||
|
)
|
||||||
|
return
|
||||||
|
assign_req_to_token_pool[(batch_size,)](
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token,
|
||||||
|
start_offset,
|
||||||
|
end_offset,
|
||||||
|
out_cache_loc,
|
||||||
|
req_to_token.shape[1],
|
||||||
|
next_power_of_2(batch_size),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _alloc_paged_token_slots_extend_npu(*args, **kwargs):
|
||||||
|
from sglang.srt.hardware_backend.npu.dsv4.dsv4_allocator import (
|
||||||
|
alloc_paged_token_slots_extend_npu,
|
||||||
|
)
|
||||||
|
|
||||||
|
return alloc_paged_token_slots_extend_npu(*args, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
ALLOC_EXTEND_FUNCS = defaultdict(
|
||||||
|
lambda: alloc_paged_token_slots_extend,
|
||||||
|
{"npu": _alloc_paged_token_slots_extend_npu},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def alloc_for_spec_decode(
|
||||||
|
tree_cache: BasePrefixCache,
|
||||||
|
req_to_token_pool: ReqToTokenPool,
|
||||||
|
*,
|
||||||
|
reqs: list[Req],
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
cur_kv_lens: torch.Tensor,
|
||||||
|
cur_kv_lens_cpu: torch.Tensor,
|
||||||
|
nxt_kv_lens: torch.Tensor,
|
||||||
|
nxt_kv_lens_cpu: torch.Tensor,
|
||||||
|
num_needed_tokens: int,
|
||||||
|
batch: Optional[ScheduleBatch] = None,
|
||||||
|
) -> None:
|
||||||
|
if num_needed_tokens > 0:
|
||||||
|
if tree_cache.token_to_kv_pool_allocator.page_size == 1:
|
||||||
|
out_cache_loc = alloc_token_slots(tree_cache, num_needed_tokens)
|
||||||
|
else:
|
||||||
|
last_loc = get_last_loc(
|
||||||
|
req_to_token_pool.req_to_token, req_pool_indices, cur_kv_lens
|
||||||
|
)
|
||||||
|
device_type = getattr(
|
||||||
|
batch.device, "type", str(batch.device).split(":", 1)[0]
|
||||||
|
)
|
||||||
|
out_cache_loc = ALLOC_EXTEND_FUNCS[device_type](
|
||||||
|
tree_cache,
|
||||||
|
cur_kv_lens,
|
||||||
|
cur_kv_lens_cpu,
|
||||||
|
nxt_kv_lens,
|
||||||
|
nxt_kv_lens_cpu,
|
||||||
|
last_loc,
|
||||||
|
num_needed_tokens,
|
||||||
|
req_pool_indices=req_pool_indices,
|
||||||
|
batch=batch,
|
||||||
|
)
|
||||||
|
# Updating req_to_token is a write to a shared tensor: it must not overlap
|
||||||
|
# with the previous batch's forward, which also reads req_to_token.
|
||||||
|
assign_req_to_token_pool_func(
|
||||||
|
req_pool_indices,
|
||||||
|
req_to_token_pool.req_to_token,
|
||||||
|
cur_kv_lens,
|
||||||
|
nxt_kv_lens,
|
||||||
|
out_cache_loc,
|
||||||
|
len(reqs),
|
||||||
|
)
|
||||||
|
|
||||||
|
for i, req in enumerate(reqs):
|
||||||
|
req.kv.kv_allocated_len = max(req.kv.kv_allocated_len, int(nxt_kv_lens_cpu[i]))
|
||||||
|
|||||||
@@ -158,7 +158,25 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
|
|||||||
return
|
return
|
||||||
|
|
||||||
start_p, end_p = effective_kv_committed_len, req.kv.kv_allocated_len
|
start_p, end_p = effective_kv_committed_len, req.kv.kv_allocated_len
|
||||||
|
_release_overallocated_kv_indices(req, start_p, end_p, tree_cache)
|
||||||
|
|
||||||
|
# If the prefix cache doesn't manage mamba states, we must free them here.
|
||||||
|
if isinstance(tree_cache.req_to_token_pool, HybridReqToTokenPool) and (
|
||||||
|
not tree_cache.supports_mamba()
|
||||||
|
):
|
||||||
|
assert (
|
||||||
|
req.mamba_pool_idx is not None
|
||||||
|
), "mamba state is freed while the tree cache does not manage mamba states"
|
||||||
|
tree_cache.req_to_token_pool.free_mamba_cache(req)
|
||||||
|
# The DSV4-NPU ReqToTokenPool subclass's free() additionally releases the
|
||||||
|
# c4/c128 state pages; other ReqToTokenPool subclasses are a no-op here.
|
||||||
|
tree_cache.req_to_token_pool.free(req)
|
||||||
|
req.kv = None
|
||||||
|
|
||||||
|
|
||||||
|
def _release_overallocated_kv_indices(
|
||||||
|
req: Req, start_p: int, end_p: int, tree_cache: BasePrefixCache
|
||||||
|
) -> None:
|
||||||
global_server_args = get_server_args()
|
global_server_args = get_server_args()
|
||||||
page_size = global_server_args.page_size
|
page_size = global_server_args.page_size
|
||||||
spec_algo = global_server_args.speculative_algorithm
|
spec_algo = global_server_args.speculative_algorithm
|
||||||
@@ -178,17 +196,6 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
|
|||||||
start_p:end_p
|
start_p:end_p
|
||||||
]
|
]
|
||||||
tree_cache.token_to_kv_pool_allocator.free(indices_to_free)
|
tree_cache.token_to_kv_pool_allocator.free(indices_to_free)
|
||||||
# If the prefix cache doesn't manage mamba states, we must free them here.
|
|
||||||
if isinstance(tree_cache.req_to_token_pool, HybridReqToTokenPool) and (
|
|
||||||
not tree_cache.supports_mamba()
|
|
||||||
):
|
|
||||||
assert (
|
|
||||||
req.mamba_pool_idx is not None
|
|
||||||
), "mamba state is freed while the tree cache does not manage mamba states"
|
|
||||||
tree_cache.req_to_token_pool.free_mamba_cache(req)
|
|
||||||
# DSV4-NPU's free() also releases c4/c128 state pages; no-op for others.
|
|
||||||
tree_cache.req_to_token_pool.free(req)
|
|
||||||
req.kv = None
|
|
||||||
|
|
||||||
|
|
||||||
def available_and_evictable_str(tree_cache: BasePrefixCache) -> str:
|
def available_and_evictable_str(tree_cache: BasePrefixCache) -> str:
|
||||||
|
|||||||
@@ -8,14 +8,9 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.mem_cache.allocation import (
|
from sglang.srt.mem_cache.allocation import alloc_for_spec_decode
|
||||||
alloc_paged_token_slots_extend,
|
|
||||||
alloc_token_slots,
|
|
||||||
get_last_loc,
|
|
||||||
)
|
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_server_args
|
||||||
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
|
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
|
||||||
from sglang.srt.speculative.spec_utils import assign_req_to_token_pool_func
|
|
||||||
from sglang.srt.utils.common import is_pin_memory_available
|
from sglang.srt.utils.common import is_pin_memory_available
|
||||||
|
|
||||||
_OVERLAP_PLAN_STREAMS: dict[str, torch.cuda.Stream] = {}
|
_OVERLAP_PLAN_STREAMS: dict[str, torch.cuda.Stream] = {}
|
||||||
@@ -193,36 +188,17 @@ class DFlashDraftInputV2(SpecInput):
|
|||||||
cur_kv_lens.copy_(cur_kv_lens_cpu_t, non_blocking=True)
|
cur_kv_lens.copy_(cur_kv_lens_cpu_t, non_blocking=True)
|
||||||
nxt_kv_lens.copy_(nxt_kv_lens_cpu_t, non_blocking=True)
|
nxt_kv_lens.copy_(nxt_kv_lens_cpu_t, non_blocking=True)
|
||||||
|
|
||||||
if num_needed_tokens > 0:
|
alloc_for_spec_decode(
|
||||||
if page_size == 1:
|
|
||||||
out_cache_loc = alloc_token_slots(
|
|
||||||
batch.tree_cache, num_needed_tokens
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
last_loc = get_last_loc(
|
|
||||||
batch.req_to_token_pool.req_to_token,
|
|
||||||
batch.req_pool_indices,
|
|
||||||
cur_kv_lens,
|
|
||||||
)
|
|
||||||
out_cache_loc = alloc_paged_token_slots_extend(
|
|
||||||
batch.tree_cache,
|
batch.tree_cache,
|
||||||
cur_kv_lens,
|
batch.req_to_token_pool,
|
||||||
cur_kv_lens_cpu_t,
|
reqs=batch.reqs,
|
||||||
nxt_kv_lens,
|
req_pool_indices=batch.req_pool_indices,
|
||||||
nxt_kv_lens_cpu_t,
|
cur_kv_lens=cur_kv_lens,
|
||||||
last_loc,
|
cur_kv_lens_cpu=cur_kv_lens_cpu_t,
|
||||||
num_needed_tokens,
|
nxt_kv_lens=nxt_kv_lens,
|
||||||
)
|
nxt_kv_lens_cpu=nxt_kv_lens_cpu_t,
|
||||||
|
num_needed_tokens=num_needed_tokens,
|
||||||
# Updating req_to_token is a write to a shared tensor: it must not overlap
|
batch=batch,
|
||||||
# with the previous batch's forward, which also reads req_to_token.
|
|
||||||
assign_req_to_token_pool_func(
|
|
||||||
batch.req_pool_indices,
|
|
||||||
batch.req_to_token_pool.req_to_token,
|
|
||||||
cur_kv_lens,
|
|
||||||
nxt_kv_lens,
|
|
||||||
out_cache_loc,
|
|
||||||
bs,
|
|
||||||
)
|
)
|
||||||
if caller_stream is not None:
|
if caller_stream is not None:
|
||||||
# Enqueue the dependency on the caller's stream, not inside the
|
# Enqueue the dependency on the caller's stream, not inside the
|
||||||
@@ -230,13 +206,6 @@ class DFlashDraftInputV2(SpecInput):
|
|||||||
# prepared req_to_token / KV allocation state.
|
# prepared req_to_token / KV allocation state.
|
||||||
caller_stream.wait_stream(plan_stream)
|
caller_stream.wait_stream(plan_stream)
|
||||||
|
|
||||||
# This request-side high-water mark is what release_kv_cache() uses to
|
|
||||||
# reclaim any DFLASH over-allocation if the request finishes later.
|
|
||||||
for i, req in enumerate(batch.reqs):
|
|
||||||
req.kv.kv_allocated_len = max(
|
|
||||||
req.kv.kv_allocated_len, int(nxt_kv_lens_cpu_t[i])
|
|
||||||
)
|
|
||||||
|
|
||||||
# Seed committed; overlap's resolve overwrites it with the published value.
|
# Seed committed; overlap's resolve overwrites it with the published value.
|
||||||
batch.seq_lens_cpu = batch_seq_lens_cpu_t
|
batch.seq_lens_cpu = batch_seq_lens_cpu_t
|
||||||
batch.seq_lens_sum = committed_seq_lens_sum
|
batch.seq_lens_sum = committed_seq_lens_sum
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
from collections import defaultdict
|
|
||||||
from enum import IntEnum
|
from enum import IntEnum
|
||||||
from typing import TYPE_CHECKING, List, Optional
|
from typing import TYPE_CHECKING, List, Optional
|
||||||
|
|
||||||
@@ -12,17 +11,10 @@ from sglang.kernels.ops.speculative.spec_tree import (
|
|||||||
sgl_build_tree_kernel_efficient_triton,
|
sgl_build_tree_kernel_efficient_triton,
|
||||||
verify_tree_greedy_kernel_triton,
|
verify_tree_greedy_kernel_triton,
|
||||||
)
|
)
|
||||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_allocator import (
|
|
||||||
alloc_paged_token_slots_extend_npu,
|
|
||||||
)
|
|
||||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
|
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
|
||||||
maybe_build_dsv4_verify_bundle,
|
maybe_build_dsv4_verify_bundle,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocation import (
|
from sglang.srt.mem_cache.allocation import alloc_for_spec_decode
|
||||||
alloc_paged_token_slots_extend,
|
|
||||||
alloc_token_slots,
|
|
||||||
get_last_loc,
|
|
||||||
)
|
|
||||||
from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode
|
from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
@@ -63,14 +55,6 @@ elif _is_cpu:
|
|||||||
from sgl_kernel import verify_tree_greedy_cpu as sgl_verify_tree_greedy_cpu
|
from sgl_kernel import verify_tree_greedy_cpu as sgl_verify_tree_greedy_cpu
|
||||||
|
|
||||||
|
|
||||||
ALLOC_EXTEND_FUNCS = defaultdict(
|
|
||||||
lambda: alloc_paged_token_slots_extend,
|
|
||||||
{
|
|
||||||
"npu": alloc_paged_token_slots_extend_npu,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def per_step_draft_out_cache_loc(
|
def per_step_draft_out_cache_loc(
|
||||||
out_cache_loc: torch.Tensor,
|
out_cache_loc: torch.Tensor,
|
||||||
batch_size: int,
|
batch_size: int,
|
||||||
@@ -817,8 +801,6 @@ def eagle_sample(
|
|||||||
def eagle_prepare_for_decode(batch: ScheduleBatch):
|
def eagle_prepare_for_decode(batch: ScheduleBatch):
|
||||||
batch.maybe_evict_swa()
|
batch.maybe_evict_swa()
|
||||||
|
|
||||||
from sglang.srt.speculative.spec_utils import assign_req_to_token_pool_func
|
|
||||||
|
|
||||||
bs = batch.batch_size()
|
bs = batch.batch_size()
|
||||||
|
|
||||||
# Accumulate penalty
|
# Accumulate penalty
|
||||||
@@ -841,7 +823,6 @@ def eagle_prepare_for_decode(batch: ScheduleBatch):
|
|||||||
cur_kv_lens[i] = cur
|
cur_kv_lens[i] = cur
|
||||||
nxt_kv_lens[i] = nxt
|
nxt_kv_lens[i] = nxt
|
||||||
num_needed_tokens += nxt - cur
|
num_needed_tokens += nxt - cur
|
||||||
r.kv.kv_allocated_len = nxt
|
|
||||||
r.decode_batch_idx += 1
|
r.decode_batch_idx += 1
|
||||||
|
|
||||||
cur_kv_lens_cpu = torch.tensor(cur_kv_lens, dtype=torch.int32, device="cpu")
|
cur_kv_lens_cpu = torch.tensor(cur_kv_lens, dtype=torch.int32, device="cpu")
|
||||||
@@ -866,31 +847,21 @@ def eagle_prepare_for_decode(batch: ScheduleBatch):
|
|||||||
# barrier has chained to the prev forward -> host stalls a full forward.
|
# barrier has chained to the prev forward -> host stalls a full forward.
|
||||||
cur_kv_lens_device = cur_kv_lens_cpu.to(device=batch.device, non_blocking=True)
|
cur_kv_lens_device = cur_kv_lens_cpu.to(device=batch.device, non_blocking=True)
|
||||||
nxt_kv_lens_device = nxt_kv_lens_cpu.to(device=batch.device, non_blocking=True)
|
nxt_kv_lens_device = nxt_kv_lens_cpu.to(device=batch.device, non_blocking=True)
|
||||||
if page_size == 1:
|
tree_cache = batch.tree_cache
|
||||||
out_cache_loc = alloc_token_slots(batch.tree_cache, num_needed_tokens)
|
req_to_token_pool = batch.req_to_token_pool
|
||||||
else:
|
req_pool_indices = batch.req_pool_indices
|
||||||
last_loc = get_last_loc(
|
reqs = batch.reqs
|
||||||
batch.req_to_token_pool.req_to_token,
|
cur_kv_lens = cur_kv_lens_device
|
||||||
batch.req_pool_indices,
|
nxt_kv_lens = nxt_kv_lens_device
|
||||||
cur_kv_lens_device,
|
alloc_for_spec_decode(
|
||||||
)
|
tree_cache,
|
||||||
device_type = getattr(batch.device, "type", str(batch.device).split(":", 1)[0])
|
req_to_token_pool,
|
||||||
out_cache_loc = ALLOC_EXTEND_FUNCS[device_type](
|
reqs=reqs,
|
||||||
batch.tree_cache,
|
req_pool_indices=req_pool_indices,
|
||||||
cur_kv_lens_device,
|
cur_kv_lens=cur_kv_lens,
|
||||||
cur_kv_lens_cpu,
|
cur_kv_lens_cpu=cur_kv_lens_cpu,
|
||||||
nxt_kv_lens_device,
|
nxt_kv_lens=nxt_kv_lens,
|
||||||
nxt_kv_lens_cpu,
|
nxt_kv_lens_cpu=nxt_kv_lens_cpu,
|
||||||
last_loc,
|
num_needed_tokens=num_needed_tokens,
|
||||||
num_needed_tokens,
|
|
||||||
req_pool_indices=batch.req_pool_indices,
|
|
||||||
batch=batch,
|
batch=batch,
|
||||||
)
|
)
|
||||||
assign_req_to_token_pool_func(
|
|
||||||
batch.req_pool_indices,
|
|
||||||
batch.req_to_token_pool.req_to_token,
|
|
||||||
cur_kv_lens_device,
|
|
||||||
nxt_kv_lens_device,
|
|
||||||
out_cache_loc,
|
|
||||||
bs,
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -16,12 +16,6 @@ from sglang.kernels.ops.speculative.cache_locs import (
|
|||||||
from sglang.kernels.ops.speculative.cache_locs import (
|
from sglang.kernels.ops.speculative.cache_locs import (
|
||||||
assign_extend_cache_locs as assign_extend_cache_locs,
|
assign_extend_cache_locs as assign_extend_cache_locs,
|
||||||
)
|
)
|
||||||
from sglang.kernels.ops.speculative.cache_locs import (
|
|
||||||
assign_req_to_token_pool as assign_req_to_token_pool,
|
|
||||||
)
|
|
||||||
from sglang.kernels.ops.speculative.cache_locs import (
|
|
||||||
assign_req_to_token_pool_func as assign_req_to_token_pool_func,
|
|
||||||
)
|
|
||||||
from sglang.kernels.ops.speculative.cache_locs import (
|
from sglang.kernels.ops.speculative.cache_locs import (
|
||||||
filter_finished_cache_loc_kernel as filter_finished_cache_loc_kernel,
|
filter_finished_cache_loc_kernel as filter_finished_cache_loc_kernel,
|
||||||
)
|
)
|
||||||
@@ -44,6 +38,12 @@ from sglang.srt.distributed.parallel_state import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.schedule_batch import set_mamba_track_indices_from_reqs
|
from sglang.srt.managers.schedule_batch import set_mamba_track_indices_from_reqs
|
||||||
|
from sglang.srt.mem_cache.allocation import (
|
||||||
|
assign_req_to_token_pool as assign_req_to_token_pool,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.allocation import (
|
||||||
|
assign_req_to_token_pool_func as assign_req_to_token_pool_func,
|
||||||
|
)
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_server_args
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
is_cpu,
|
is_cpu,
|
||||||
|
|||||||
@@ -85,6 +85,7 @@ class TestDeepSeekV4HiSparseAllocator(CustomTestCase):
|
|||||||
rid="req-0",
|
rid="req-0",
|
||||||
origin_input_ids=list(range(fill_len)),
|
origin_input_ids=list(range(fill_len)),
|
||||||
output_ids=[],
|
output_ids=[],
|
||||||
|
kv=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
def set_extend_range(start, end):
|
def set_extend_range(start, end):
|
||||||
@@ -143,8 +144,8 @@ class TestDeepSeekV4HiSparseAllocator(CustomTestCase):
|
|||||||
_, kwargs = allocator.alloc_extend_swa_tail.call_args
|
_, kwargs = allocator.alloc_extend_swa_tail.call_args
|
||||||
self.assertEqual(kwargs["extend_num_tokens"], fill_len)
|
self.assertEqual(kwargs["extend_num_tokens"], fill_len)
|
||||||
self.assertEqual(kwargs["swa_tail_len"], swa_tail_len)
|
self.assertEqual(kwargs["swa_tail_len"], swa_tail_len)
|
||||||
self.assertEqual(req.swa_evicted_seqlen, fill_len - swa_tail_len)
|
self.assertEqual(req.kv.swa_evicted_seqlen, fill_len - swa_tail_len)
|
||||||
self.assertEqual(req.kv_allocated_len, fill_len)
|
self.assertEqual(req.kv.kv_allocated_len, fill_len)
|
||||||
self.assertEqual(req.kv_committed_len, fill_len)
|
self.assertEqual(req.kv_committed_len, fill_len)
|
||||||
self.assertEqual(req.extend_range.length, fill_len)
|
self.assertEqual(req.extend_range.length, fill_len)
|
||||||
self.assertEqual(len(req_to_token_pool.writes), 1)
|
self.assertEqual(len(req_to_token_pool.writes), 1)
|
||||||
|
|||||||
@@ -4071,7 +4071,7 @@ class TestUnifiedRadixCacheInt8MambaCheckpoint(CustomTestCase):
|
|||||||
req_to_token_pool.alloc([req])
|
req_to_token_pool.alloc([req])
|
||||||
req.output_ids = array("q")
|
req.output_ids = array("q")
|
||||||
req.kv_committed_len = len(tokens)
|
req.kv_committed_len = len(tokens)
|
||||||
req.kv_allocated_len = len(tokens)
|
req.kv = ReqKvInfo(kv_allocated_len=len(tokens), swa_evicted_seqlen=0)
|
||||||
req.cache_protected_len = 0
|
req.cache_protected_len = 0
|
||||||
req.swa_uuid_for_lock = None
|
req.swa_uuid_for_lock = None
|
||||||
req.extra_key = None
|
req.extra_key = None
|
||||||
@@ -4085,7 +4085,9 @@ class TestUnifiedRadixCacheInt8MambaCheckpoint(CustomTestCase):
|
|||||||
req_to_token_pool.write((req.req_pool_idx, slice(0, len(tokens))), kv_indices)
|
req_to_token_pool.write((req.req_pool_idx, slice(0, len(tokens))), kv_indices)
|
||||||
req.last_node = cache.root_node
|
req.last_node = cache.root_node
|
||||||
|
|
||||||
cache.cache_finished_req(req, is_insert=True)
|
cache.cache_finished_req(
|
||||||
|
req, is_insert=True, kv_len_to_handle=req.effective_kv_committed_len()
|
||||||
|
)
|
||||||
|
|
||||||
def test_finished_req_stores_radix_mamba_state_in_int8_pool(self):
|
def test_finished_req_stores_radix_mamba_state_in_int8_pool(self):
|
||||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||||
|
|||||||
@@ -50,32 +50,36 @@ _OWNER_SITES = {
|
|||||||
# non-spec scheduler
|
# non-spec scheduler
|
||||||
(_SB, "ScheduleBatch.prepare_for_decode", "decode_batch_idx"): 1,
|
(_SB, "ScheduleBatch.prepare_for_decode", "decode_batch_idx"): 1,
|
||||||
(_SB, "ScheduleBatch.prepare_for_decode", "kv_committed_len"): 1,
|
(_SB, "ScheduleBatch.prepare_for_decode", "kv_committed_len"): 1,
|
||||||
(_SB, "ScheduleBatch.prepare_for_decode", "kv_allocated_len"): 1,
|
|
||||||
(_SB, "ScheduleBatch.prepare_for_extend", "extend_batch_idx"): 1,
|
(_SB, "ScheduleBatch.prepare_for_extend", "extend_batch_idx"): 1,
|
||||||
(_SB, "ScheduleBatch.prepare_for_extend", "kv_committed_len"): 1,
|
(_SB, "ScheduleBatch.prepare_for_extend", "kv_committed_len"): 1,
|
||||||
(_SB, "ScheduleBatch.prepare_for_extend", "kv_allocated_len"): 1,
|
# kv_allocated_len is settled inside the owned-kv alloc functions (op28).
|
||||||
("mem_cache/allocation.py", "alloc_for_extend", "evict"): 1,
|
("mem_cache/allocation.py", "alloc_for_extend", "evict"): 1,
|
||||||
|
("mem_cache/allocation.py", "alloc_for_extend", "kv_allocated_len"): 1,
|
||||||
("mem_cache/allocation.py", "alloc_for_decode", "evict"): 1,
|
("mem_cache/allocation.py", "alloc_for_decode", "evict"): 1,
|
||||||
|
("mem_cache/allocation.py", "alloc_for_decode", "kv_allocated_len"): 1,
|
||||||
# spec v2: no pre-claim; resolve commits the full accepted run uniformly.
|
# spec v2: no pre-claim; resolve commits the full accepted run uniformly.
|
||||||
|
# kv_allocated_len for spec v2 draft decode (eagle + dflash) is settled
|
||||||
|
# inside the owned-kv alloc_for_spec_decode function (op42).
|
||||||
(*_EAGLE_DECODE, "decode_batch_idx"): 1,
|
(*_EAGLE_DECODE, "decode_batch_idx"): 1,
|
||||||
(*_EAGLE_DECODE, "evict"): 1,
|
(*_EAGLE_DECODE, "evict"): 1,
|
||||||
(*_EAGLE_DECODE, "kv_allocated_len"): 1,
|
|
||||||
(*_RESOLVE, "kv_committed_len"): 1,
|
|
||||||
(*_RESOLVE, "spec_verify_ct"): 1,
|
|
||||||
(
|
(
|
||||||
"speculative/dflash_info_v2.py",
|
"mem_cache/allocation.py",
|
||||||
"DFlashDraftInputV2.prepare_for_decode",
|
"alloc_for_spec_decode",
|
||||||
"kv_allocated_len",
|
"kv_allocated_len",
|
||||||
): 1,
|
): 1,
|
||||||
# disaggregation decode prealloc
|
(*_RESOLVE, "kv_committed_len"): 1,
|
||||||
|
(*_RESOLVE, "spec_verify_ct"): 1,
|
||||||
|
# disaggregation decode prealloc: kv_allocated_len is settled inside the
|
||||||
|
# owned-kv alloc_for_decode_prealloc(_hisparse) functions (op42).
|
||||||
(
|
(
|
||||||
"disaggregation/decode.py",
|
"disaggregation/decode.py",
|
||||||
"DecodePreallocQueue._pre_alloc",
|
"DecodePreallocQueue._pre_alloc",
|
||||||
"kv_committed_len",
|
"kv_committed_len",
|
||||||
): 1,
|
): 1,
|
||||||
|
("disaggregation/decode.py", "alloc_for_decode_prealloc", "kv_allocated_len"): 1,
|
||||||
(
|
(
|
||||||
"disaggregation/decode.py",
|
"disaggregation/decode.py",
|
||||||
"DecodePreallocQueue._pre_alloc",
|
"alloc_for_decode_prealloc_hisparse",
|
||||||
"kv_allocated_len",
|
"kv_allocated_len",
|
||||||
): 1,
|
): 1,
|
||||||
# streaming session slot save/restore and tail trimming
|
# streaming session slot save/restore and tail trimming
|
||||||
|
|||||||
Reference in New Issue
Block a user