[dLLM] Reuse block KV/req slots in place across FDFO rounds (#27877)
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
co-authored by
Xiaoyu Zhang
parent
b15a83983c
commit
49b9c46f41
@@ -78,8 +78,7 @@ class SchedulerDllmMixin:
|
|||||||
not fdfo_mode or result.accept_length_per_req_cpu is not None
|
not fdfo_mode or result.accept_length_per_req_cpu is not None
|
||||||
), "FDFO dLLM result is missing accept lengths."
|
), "FDFO dLLM result is missing accept lengths."
|
||||||
|
|
||||||
# Sync mode emits tokens only once a block fully resolves; FDFO always
|
# FDFO also commits unresolved blocks so their KV can be reused.
|
||||||
# commits (resolved blocks decode, unresolved blocks stash + free KV).
|
|
||||||
if fdfo_mode or result.next_token_ids:
|
if fdfo_mode or result.next_token_ids:
|
||||||
block_size = self.dllm_config.block_size
|
block_size = self.dllm_config.block_size
|
||||||
algo_states = result.dllm_algo_state
|
algo_states = result.dllm_algo_state
|
||||||
@@ -111,20 +110,11 @@ class SchedulerDllmMixin:
|
|||||||
assert len(next_token_ids) == block_size
|
assert len(next_token_ids) == block_size
|
||||||
|
|
||||||
if result.accept_length_per_req_cpu[idx] == 0:
|
if result.accept_length_per_req_cpu[idx] == 0:
|
||||||
# Block unresolved: stash partial state and free the KV slots
|
# Unresolved: keep partial state and KV for the next FDFO round.
|
||||||
# of the still-masked block so the next FDFO round can
|
|
||||||
# re-denoise it without leaking the previous allocation.
|
|
||||||
req.dllm_incomplete_ids = array("q", next_token_ids)
|
req.dllm_incomplete_ids = array("q", next_token_ids)
|
||||||
req.dllm_algo_state = (
|
req.dllm_algo_state = (
|
||||||
algo_states[idx] if algo_states is not None else None
|
algo_states[idx] if algo_states is not None else None
|
||||||
)
|
)
|
||||||
old_prefix_len = len(req.prefix_indices)
|
|
||||||
new_fill_len = req.extend_range.end
|
|
||||||
if new_fill_len > old_prefix_len:
|
|
||||||
kv_indices_to_free = self.req_to_token_pool.req_to_token[
|
|
||||||
req.req_pool_idx, old_prefix_len:new_fill_len
|
|
||||||
]
|
|
||||||
self.token_to_kv_pool_allocator.free(kv_indices_to_free)
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
req.dllm_incomplete_ids = array("q")
|
req.dllm_incomplete_ids = array("q")
|
||||||
@@ -426,6 +416,25 @@ class DllmManager:
|
|||||||
self.waiting_queue = [req for req in self.waiting_queue if not req.finished()]
|
self.waiting_queue = [req for req in self.waiting_queue if not req.finished()]
|
||||||
self.staging_queue = [req for req in self.staging_queue if not req.finished()]
|
self.staging_queue = [req for req in self.staging_queue if not req.finished()]
|
||||||
|
|
||||||
|
def pop_aborted_reqs(self, abort_all: bool, rid: str) -> List[Req]:
|
||||||
|
aborted_reqs: List[Req] = []
|
||||||
|
seen: Set[int] = set()
|
||||||
|
|
||||||
|
for queue_name in ("waiting_queue", "staging_queue"):
|
||||||
|
queue = getattr(self, queue_name)
|
||||||
|
kept_queue = []
|
||||||
|
for req in queue:
|
||||||
|
if abort_all or req.rid.startswith(rid):
|
||||||
|
req_id = id(req)
|
||||||
|
if req_id not in seen:
|
||||||
|
aborted_reqs.append(req)
|
||||||
|
seen.add(req_id)
|
||||||
|
else:
|
||||||
|
kept_queue.append(req)
|
||||||
|
setattr(self, queue_name, kept_queue)
|
||||||
|
|
||||||
|
return aborted_reqs
|
||||||
|
|
||||||
def init_next_round(self) -> None:
|
def init_next_round(self) -> None:
|
||||||
"""Initialize staging requests for next round and clear staging queue."""
|
"""Initialize staging requests for next round and clear staging queue."""
|
||||||
for req in self.staging_queue:
|
for req in self.staging_queue:
|
||||||
|
|||||||
@@ -776,6 +776,8 @@ class PrefillAdder:
|
|||||||
cand_extend_input_len = len(req.full_untruncated_fill_ids) - len(
|
cand_extend_input_len = len(req.full_untruncated_fill_ids) - len(
|
||||||
req.prefix_indices
|
req.prefix_indices
|
||||||
)
|
)
|
||||||
|
if req.dllm_incomplete_ids and cand_extend_input_len > _rem_tokens:
|
||||||
|
return AddReqResult.NO_TOKEN
|
||||||
truncated = cand_extend_input_len > _rem_tokens
|
truncated = cand_extend_input_len > _rem_tokens
|
||||||
new_len = min(cand_extend_input_len, _rem_tokens)
|
new_len = min(cand_extend_input_len, _rem_tokens)
|
||||||
req.set_extend_range(len(req.prefix_indices), len(req.prefix_indices) + new_len)
|
req.set_extend_range(len(req.prefix_indices), len(req.prefix_indices) + new_len)
|
||||||
|
|||||||
@@ -2693,7 +2693,8 @@ class Scheduler(
|
|||||||
if self.dllm_config.first_done_first_out_mode:
|
if self.dllm_config.first_done_first_out_mode:
|
||||||
if not req.dllm_incomplete_ids:
|
if not req.dllm_incomplete_ids:
|
||||||
self.stash_chunked_request(req)
|
self.stash_chunked_request(req)
|
||||||
self.req_to_token_pool.free(req)
|
self.req_to_token_pool.free(req)
|
||||||
|
# Otherwise, keep req slot/KV for reuse.
|
||||||
else:
|
else:
|
||||||
self.stash_chunked_request(req)
|
self.stash_chunked_request(req)
|
||||||
|
|
||||||
@@ -4061,6 +4062,22 @@ class Scheduler(
|
|||||||
release_kv_cache(req, self.tree_cache, is_insert=False)
|
release_kv_cache(req, self.tree_cache, is_insert=False)
|
||||||
logger.debug(f"Abort queued request. {req.rid=}")
|
logger.debug(f"Abort queued request. {req.rid=}")
|
||||||
|
|
||||||
|
if self.dllm_config is not None:
|
||||||
|
for req in self.dllm_manager.pop_aborted_reqs(
|
||||||
|
recv_req.abort_all, recv_req.rid
|
||||||
|
):
|
||||||
|
if self.enable_hicache_storage:
|
||||||
|
self.tree_cache.release_aborted_request(req.rid)
|
||||||
|
self.ipc_channels.send_to_tokenizer.send_output(
|
||||||
|
AbortReq(rid=req.rid), req
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
req.req_pool_idx is not None
|
||||||
|
or getattr(req, "mamba_pool_idx", None) is not None
|
||||||
|
):
|
||||||
|
release_kv_cache(req, self.tree_cache, is_insert=False)
|
||||||
|
logger.debug(f"Abort dLLM queued request. {req.rid=}")
|
||||||
|
|
||||||
# Delete the requests in the grammar queue
|
# Delete the requests in the grammar queue
|
||||||
# Abort method 2: call `set_finish_with_abort`
|
# Abort method 2: call `set_finish_with_abort`
|
||||||
# The request will still run one prefill forward pass.
|
# The request will still run one prefill forward pass.
|
||||||
|
|||||||
@@ -315,6 +315,13 @@ def alloc_for_extend(
|
|||||||
|
|
||||||
prefix_tensors = [r.prefix_indices for r in batch.reqs]
|
prefix_tensors = [r.prefix_indices for r in batch.reqs]
|
||||||
|
|
||||||
|
reuse_kv = None
|
||||||
|
if batch.is_dllm():
|
||||||
|
reuse_kv = [
|
||||||
|
r.req_pool_idx is not None and bool(r.dllm_incomplete_ids)
|
||||||
|
for r in batch.reqs
|
||||||
|
]
|
||||||
|
|
||||||
# Create tensors for allocation
|
# Create tensors for allocation
|
||||||
prefix_lens_cpu = torch.tensor(batch.prefix_lens, dtype=torch.int64)
|
prefix_lens_cpu = torch.tensor(batch.prefix_lens, dtype=torch.int64)
|
||||||
extend_lens_cpu = torch.tensor(batch.extend_lens, dtype=torch.int64)
|
extend_lens_cpu = torch.tensor(batch.extend_lens, dtype=torch.int64)
|
||||||
@@ -329,7 +336,18 @@ def alloc_for_extend(
|
|||||||
req_pool_indices_device = req_pool_indices_cpu.to(batch.device, non_blocking=True)
|
req_pool_indices_device = req_pool_indices_cpu.to(batch.device, non_blocking=True)
|
||||||
|
|
||||||
# Allocate KV cache (throws exception on failure)
|
# Allocate KV cache (throws exception on failure)
|
||||||
if _alloc_page_size(batch) == 1:
|
alloc_page_size = _alloc_page_size(batch)
|
||||||
|
if reuse_kv is not None and any(reuse_kv):
|
||||||
|
out_cache_loc = _alloc_extend_loc_with_kv_reuse(
|
||||||
|
batch,
|
||||||
|
reuse_kv,
|
||||||
|
req_pool_indices_cpu,
|
||||||
|
prefix_lens_cpu,
|
||||||
|
extend_lens_cpu,
|
||||||
|
req_pool_indices_device,
|
||||||
|
alloc_page_size,
|
||||||
|
)
|
||||||
|
elif alloc_page_size == 1:
|
||||||
out_cache_loc = alloc_token_slots(batch.tree_cache, batch.extend_num_tokens)
|
out_cache_loc = alloc_token_slots(batch.tree_cache, batch.extend_num_tokens)
|
||||||
else:
|
else:
|
||||||
# Paged allocation - build last_loc
|
# Paged allocation - build last_loc
|
||||||
@@ -385,6 +403,87 @@ def alloc_for_extend(
|
|||||||
return out_cache_loc, req_pool_indices_device, req_pool_indices_cpu
|
return out_cache_loc, req_pool_indices_device, req_pool_indices_cpu
|
||||||
|
|
||||||
|
|
||||||
|
def _alloc_extend_loc_with_kv_reuse(
|
||||||
|
batch: ScheduleBatch,
|
||||||
|
reuse_kv: list[bool],
|
||||||
|
req_pool_indices_cpu: torch.Tensor,
|
||||||
|
prefix_lens_cpu: torch.Tensor,
|
||||||
|
extend_lens_cpu: torch.Tensor,
|
||||||
|
req_pool_indices_device: torch.Tensor,
|
||||||
|
alloc_page_size: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
device = batch.device
|
||||||
|
req_to_token = batch.req_to_token_pool.req_to_token
|
||||||
|
|
||||||
|
for i, req in enumerate(batch.reqs):
|
||||||
|
if not reuse_kv[i]:
|
||||||
|
continue
|
||||||
|
prefix_len = int(prefix_lens_cpu[i])
|
||||||
|
extend_len = int(extend_lens_cpu[i])
|
||||||
|
retained_len = len(req.dllm_incomplete_ids)
|
||||||
|
if extend_len != retained_len:
|
||||||
|
raise RuntimeError("dLLM FDFO retained KV must be reused as a full block.")
|
||||||
|
if req.kv is None or prefix_len + extend_len > req.kv.kv_allocated_len:
|
||||||
|
raise RuntimeError("dLLM FDFO retained KV is missing.")
|
||||||
|
|
||||||
|
alloc_extend_lens = [
|
||||||
|
0 if reuse_kv[i] else int(extend_lens_cpu[i]) for i in range(len(reuse_kv))
|
||||||
|
]
|
||||||
|
alloc_extend_num_tokens = sum(alloc_extend_lens)
|
||||||
|
|
||||||
|
fresh_slots = None
|
||||||
|
if alloc_extend_num_tokens > 0:
|
||||||
|
if alloc_page_size == 1:
|
||||||
|
fresh_slots = alloc_token_slots(batch.tree_cache, alloc_extend_num_tokens)
|
||||||
|
else:
|
||||||
|
alloc_seq_lens_cpu = torch.tensor(
|
||||||
|
[
|
||||||
|
(
|
||||||
|
int(prefix_lens_cpu[i])
|
||||||
|
if reuse_kv[i]
|
||||||
|
else int(batch.seq_lens_cpu[i])
|
||||||
|
)
|
||||||
|
for i in range(len(reuse_kv))
|
||||||
|
],
|
||||||
|
dtype=torch.int64,
|
||||||
|
)
|
||||||
|
last_loc = [
|
||||||
|
(t[-1:] if len(t) > 0 else torch.tensor([-1], device=device))
|
||||||
|
for t in (r.prefix_indices for r in batch.reqs)
|
||||||
|
]
|
||||||
|
fresh_slots = alloc_paged_token_slots_extend(
|
||||||
|
tree_cache=batch.tree_cache,
|
||||||
|
prefix_lens=prefix_lens_cpu.to(device, non_blocking=True),
|
||||||
|
prefix_lens_cpu=prefix_lens_cpu,
|
||||||
|
seq_lens=alloc_seq_lens_cpu.to(device, non_blocking=True),
|
||||||
|
seq_lens_cpu=alloc_seq_lens_cpu,
|
||||||
|
last_loc=torch.cat(last_loc),
|
||||||
|
extend_num_tokens=alloc_extend_num_tokens,
|
||||||
|
req_pool_indices=req_pool_indices_device,
|
||||||
|
dsv4_state_lens=_compute_dsv4_state_lens(batch, is_decode=False),
|
||||||
|
batch=batch,
|
||||||
|
)
|
||||||
|
|
||||||
|
reuse_dtype = fresh_slots.dtype if fresh_slots is not None else torch.int64
|
||||||
|
parts: list[torch.Tensor] = []
|
||||||
|
fresh_ptr = 0
|
||||||
|
for i in range(len(reuse_kv)):
|
||||||
|
prefix_len = int(prefix_lens_cpu[i])
|
||||||
|
extend_len = int(extend_lens_cpu[i])
|
||||||
|
if reuse_kv[i]:
|
||||||
|
req_idx = int(req_pool_indices_cpu[i])
|
||||||
|
parts.append(
|
||||||
|
req_to_token[req_idx, prefix_len : prefix_len + extend_len].to(
|
||||||
|
reuse_dtype
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
parts.append(fresh_slots[fresh_ptr : fresh_ptr + extend_len])
|
||||||
|
fresh_ptr += extend_len
|
||||||
|
|
||||||
|
return torch.cat(parts)
|
||||||
|
|
||||||
|
|
||||||
def alloc_paged_token_slots_decode(
|
def alloc_paged_token_slots_decode(
|
||||||
tree_cache: BasePrefixCache,
|
tree_cache: BasePrefixCache,
|
||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
|
|||||||
@@ -0,0 +1,219 @@
|
|||||||
|
"""Tests dLLM FDFO KV slot reuse in alloc_for_extend."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from array import array
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.dllm.mixin.scheduler import DllmManager
|
||||||
|
from sglang.srt.mem_cache import allocation
|
||||||
|
from sglang.srt.mem_cache.allocation import alloc_for_extend
|
||||||
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeAllocator:
|
||||||
|
def __init__(self, base=1000, page_size=1):
|
||||||
|
self.base = base
|
||||||
|
self.page_size = page_size
|
||||||
|
self.alloc_calls = []
|
||||||
|
self.extend_calls = []
|
||||||
|
|
||||||
|
def available_size(self):
|
||||||
|
return 1 << 30
|
||||||
|
|
||||||
|
def alloc(self, need_size):
|
||||||
|
self.alloc_calls.append(need_size)
|
||||||
|
return torch.arange(self.base, self.base + need_size, dtype=torch.int64)
|
||||||
|
|
||||||
|
def alloc_extend(
|
||||||
|
self,
|
||||||
|
prefix_lens,
|
||||||
|
prefix_lens_cpu,
|
||||||
|
seq_lens,
|
||||||
|
seq_lens_cpu,
|
||||||
|
last_loc,
|
||||||
|
extend_num_tokens,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
self.extend_calls.append(
|
||||||
|
{
|
||||||
|
"extend_num_tokens": extend_num_tokens,
|
||||||
|
"seq_lens_cpu": seq_lens_cpu.tolist(),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return torch.arange(self.base, self.base + extend_num_tokens, dtype=torch.int64)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeTreeCache:
|
||||||
|
def __init__(self, allocator):
|
||||||
|
self.page_size = allocator.page_size
|
||||||
|
self.token_to_kv_pool_allocator = allocator
|
||||||
|
|
||||||
|
def is_chunk_cache(self):
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def _make_req(rid, prefix, block_size, *, req_pool_idx=None, reuse=False):
|
||||||
|
return SimpleNamespace(
|
||||||
|
rid=rid,
|
||||||
|
prefix_indices=torch.tensor(prefix, dtype=torch.int32),
|
||||||
|
req_pool_idx=req_pool_idx,
|
||||||
|
dllm_incomplete_ids=array("q", range(block_size)) if reuse else array("q"),
|
||||||
|
inflight_middle_chunks=1 if req_pool_idx is not None else 0,
|
||||||
|
kv_committed_len=len(prefix) if req_pool_idx is not None else 0,
|
||||||
|
kv=(
|
||||||
|
SimpleNamespace(kv_allocated_len=len(prefix) + block_size)
|
||||||
|
if req_pool_idx is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _remove_allocated_req_slots(pool, *reqs):
|
||||||
|
for req in reqs:
|
||||||
|
if req.req_pool_idx in pool.free_slots:
|
||||||
|
pool.free_slots.remove(req.req_pool_idx)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_batch(pool, allocator, reqs, extend_lens):
|
||||||
|
seq_lens_cpu = torch.tensor(
|
||||||
|
[
|
||||||
|
len(req.prefix_indices) + extend_len
|
||||||
|
for req, extend_len in zip(reqs, extend_lens)
|
||||||
|
],
|
||||||
|
dtype=torch.int64,
|
||||||
|
)
|
||||||
|
return SimpleNamespace(
|
||||||
|
device="cpu",
|
||||||
|
reqs=reqs,
|
||||||
|
req_to_token_pool=pool,
|
||||||
|
token_to_kv_pool_allocator=allocator,
|
||||||
|
tree_cache=_FakeTreeCache(allocator),
|
||||||
|
prefix_lens=[len(req.prefix_indices) for req in reqs],
|
||||||
|
extend_lens=extend_lens,
|
||||||
|
seq_lens=seq_lens_cpu,
|
||||||
|
seq_lens_cpu=seq_lens_cpu,
|
||||||
|
extend_num_tokens=sum(extend_lens),
|
||||||
|
maybe_evict_swa=lambda: None,
|
||||||
|
is_dllm=lambda: True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _seed_retained_block(pool, req, values):
|
||||||
|
prefix_len = len(req.prefix_indices)
|
||||||
|
pool.req_to_token[req.req_pool_idx, :prefix_len] = req.prefix_indices
|
||||||
|
pool.req_to_token[req.req_pool_idx, prefix_len : prefix_len + len(values)] = (
|
||||||
|
torch.tensor(values, dtype=torch.int32)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDllmFdfoKvReuse(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.block_size = 4
|
||||||
|
self.pool = ReqToTokenPool(
|
||||||
|
size=8, max_context_len=64, device="cpu", enable_memory_saver=False
|
||||||
|
)
|
||||||
|
self._old_support_triton = allocation.support_triton
|
||||||
|
self._old_get_server_args = allocation.get_server_args
|
||||||
|
allocation.support_triton = lambda _: False
|
||||||
|
allocation.get_server_args = lambda: SimpleNamespace(
|
||||||
|
attention_backend="torch_native", dcp_size=1
|
||||||
|
)
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
allocation.support_triton = self._old_support_triton
|
||||||
|
allocation.get_server_args = self._old_get_server_args
|
||||||
|
|
||||||
|
def test_alloc_for_extend_mixed_reuse_allocates_only_fresh_and_writes_rows(self):
|
||||||
|
allocator = _FakeAllocator(base=200)
|
||||||
|
reused = _make_req(
|
||||||
|
"reuse", [10, 11, 12, 13], self.block_size, req_pool_idx=1, reuse=True
|
||||||
|
)
|
||||||
|
fresh = _make_req("fresh", [20, 21, 22, 23], self.block_size)
|
||||||
|
_remove_allocated_req_slots(self.pool, reused)
|
||||||
|
_seed_retained_block(self.pool, reused, [100, 101, 102, 103])
|
||||||
|
|
||||||
|
batch = _make_batch(self.pool, allocator, [reused, fresh], [4, 4])
|
||||||
|
out, _, req_pool_indices_cpu = alloc_for_extend(batch)
|
||||||
|
|
||||||
|
self.assertEqual(allocator.alloc_calls, [4])
|
||||||
|
self.assertEqual(req_pool_indices_cpu.tolist(), [1, 2])
|
||||||
|
self.assertEqual(out.tolist(), [100, 101, 102, 103, 200, 201, 202, 203])
|
||||||
|
self.assertEqual(self.pool.req_to_token[1, 4:8].tolist(), [100, 101, 102, 103])
|
||||||
|
self.assertEqual(self.pool.req_to_token[2, 4:8].tolist(), [200, 201, 202, 203])
|
||||||
|
self.assertEqual(reused.kv.kv_allocated_len, 8)
|
||||||
|
self.assertEqual(fresh.kv.kv_allocated_len, 8)
|
||||||
|
|
||||||
|
def test_alloc_for_extend_all_reuse_allocates_nothing(self):
|
||||||
|
allocator = _FakeAllocator(base=900)
|
||||||
|
req0 = _make_req(
|
||||||
|
"r0", [1, 2, 3, 4], self.block_size, req_pool_idx=1, reuse=True
|
||||||
|
)
|
||||||
|
req1 = _make_req(
|
||||||
|
"r1", [5, 6, 7, 8], self.block_size, req_pool_idx=2, reuse=True
|
||||||
|
)
|
||||||
|
_remove_allocated_req_slots(self.pool, req0, req1)
|
||||||
|
_seed_retained_block(self.pool, req0, [300, 301, 302, 303])
|
||||||
|
_seed_retained_block(self.pool, req1, [400, 401, 402, 403])
|
||||||
|
|
||||||
|
batch = _make_batch(self.pool, allocator, [req0, req1], [4, 4])
|
||||||
|
out, _, req_pool_indices_cpu = alloc_for_extend(batch)
|
||||||
|
|
||||||
|
self.assertEqual(allocator.alloc_calls, [])
|
||||||
|
self.assertEqual(req_pool_indices_cpu.tolist(), [1, 2])
|
||||||
|
self.assertEqual(out.tolist(), [300, 301, 302, 303, 400, 401, 402, 403])
|
||||||
|
|
||||||
|
def test_alloc_for_extend_paged_mixed_reuse_skips_reused_rows(self):
|
||||||
|
allocator = _FakeAllocator(base=500, page_size=4)
|
||||||
|
reused = _make_req(
|
||||||
|
"reuse", [10, 11, 12, 13], self.block_size, req_pool_idx=1, reuse=True
|
||||||
|
)
|
||||||
|
fresh = _make_req("fresh", [20, 21, 22, 23], self.block_size)
|
||||||
|
_remove_allocated_req_slots(self.pool, reused)
|
||||||
|
_seed_retained_block(self.pool, reused, [100, 101, 102, 103])
|
||||||
|
|
||||||
|
batch = _make_batch(self.pool, allocator, [reused, fresh], [4, 4])
|
||||||
|
out, _, req_pool_indices_cpu = alloc_for_extend(batch)
|
||||||
|
|
||||||
|
self.assertEqual(req_pool_indices_cpu.tolist(), [1, 2])
|
||||||
|
self.assertEqual(out.tolist(), [100, 101, 102, 103, 500, 501, 502, 503])
|
||||||
|
self.assertEqual(
|
||||||
|
allocator.extend_calls,
|
||||||
|
[{"extend_num_tokens": 4, "seq_lens_cpu": [4, 8]}],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_alloc_for_extend_rejects_partial_retained_block_reuse(self):
|
||||||
|
allocator = _FakeAllocator(base=700)
|
||||||
|
reused = _make_req(
|
||||||
|
"reuse", [10, 11, 12, 13], self.block_size, req_pool_idx=1, reuse=True
|
||||||
|
)
|
||||||
|
_remove_allocated_req_slots(self.pool, reused)
|
||||||
|
_seed_retained_block(self.pool, reused, [100, 101, 102, 103])
|
||||||
|
|
||||||
|
batch = _make_batch(self.pool, allocator, [reused], [2])
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "full block"):
|
||||||
|
alloc_for_extend(batch)
|
||||||
|
|
||||||
|
def test_dllm_manager_pop_aborted_reqs_removes_waiting_and_staging(self):
|
||||||
|
manager = DllmManager(SimpleNamespace(max_running_requests=4))
|
||||||
|
waiting = _make_req("abort-waiting", [1], self.block_size)
|
||||||
|
staging = _make_req("abort-staging", [2], self.block_size)
|
||||||
|
keep = _make_req("keep", [3], self.block_size)
|
||||||
|
manager.waiting_queue = [waiting, keep]
|
||||||
|
manager.staging_queue = [staging, waiting]
|
||||||
|
|
||||||
|
aborted = manager.pop_aborted_reqs(False, "abort")
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
[req.rid for req in aborted], ["abort-waiting", "abort-staging"]
|
||||||
|
)
|
||||||
|
self.assertEqual(manager.waiting_queue, [keep])
|
||||||
|
self.assertEqual(manager.staging_queue, [])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user