diff --git a/python/sglang/srt/dllm/mixin/scheduler.py b/python/sglang/srt/dllm/mixin/scheduler.py index 2da58fb0c..6d532531e 100644 --- a/python/sglang/srt/dllm/mixin/scheduler.py +++ b/python/sglang/srt/dllm/mixin/scheduler.py @@ -78,8 +78,7 @@ class SchedulerDllmMixin: not fdfo_mode or result.accept_length_per_req_cpu is not None ), "FDFO dLLM result is missing accept lengths." - # Sync mode emits tokens only once a block fully resolves; FDFO always - # commits (resolved blocks decode, unresolved blocks stash + free KV). + # FDFO also commits unresolved blocks so their KV can be reused. if fdfo_mode or result.next_token_ids: block_size = self.dllm_config.block_size algo_states = result.dllm_algo_state @@ -111,20 +110,11 @@ class SchedulerDllmMixin: assert len(next_token_ids) == block_size if result.accept_length_per_req_cpu[idx] == 0: - # Block unresolved: stash partial state and free the KV slots - # of the still-masked block so the next FDFO round can - # re-denoise it without leaking the previous allocation. + # Unresolved: keep partial state and KV for the next FDFO round. req.dllm_incomplete_ids = array("q", next_token_ids) req.dllm_algo_state = ( 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 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.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: """Initialize staging requests for next round and clear staging queue.""" for req in self.staging_queue: diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index c1f359599..1e24895f4 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -776,6 +776,8 @@ class PrefillAdder: cand_extend_input_len = len(req.full_untruncated_fill_ids) - len( 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 new_len = min(cand_extend_input_len, _rem_tokens) req.set_extend_range(len(req.prefix_indices), len(req.prefix_indices) + new_len) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 99da61535..cb4a54768 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2693,7 +2693,8 @@ class Scheduler( if self.dllm_config.first_done_first_out_mode: if not req.dllm_incomplete_ids: 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: self.stash_chunked_request(req) @@ -4061,6 +4062,22 @@ class Scheduler( release_kv_cache(req, self.tree_cache, is_insert=False) 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 # Abort method 2: call `set_finish_with_abort` # The request will still run one prefill forward pass. diff --git a/python/sglang/srt/mem_cache/allocation.py b/python/sglang/srt/mem_cache/allocation.py index 4ac59c97d..3d6c8cf86 100644 --- a/python/sglang/srt/mem_cache/allocation.py +++ b/python/sglang/srt/mem_cache/allocation.py @@ -315,6 +315,13 @@ def alloc_for_extend( 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 prefix_lens_cpu = torch.tensor(batch.prefix_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) # 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) else: # 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 +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( tree_cache: BasePrefixCache, seq_lens: torch.Tensor, diff --git a/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py b/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py new file mode 100644 index 000000000..c48fcc0e9 --- /dev/null +++ b/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py @@ -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()