From 44e3dd2713de3c53afeb0a47b8b71293ee213e37 Mon Sep 17 00:00:00 2001 From: huangtingwei <141888744+huangtingwei9988@users.noreply.github.com> Date: Fri, 17 Jul 2026 14:52:08 +0800 Subject: [PATCH] [HiCache] Optimize HiCache host pool free-list release (#30658) Co-authored-by: Zhangheng --- .../sglang/srt/mem_cache/memory_pool_host.py | 32 +++++-- python/sglang/srt/mem_cache/pool_host/base.py | 27 +++++- .../unit/mem_cache/test_mem_pool_host.py | 87 ++++++++++++++++++- 3 files changed, 137 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 729aa8d42..139685227 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -249,9 +249,11 @@ class MambaPoolHost(HostKVCache): (self.size,), dtype=torch.uint8, device=self.device ) self.free_slots = torch.arange(self.size, dtype=torch.int64) + self.release_slots = [] + self.num_release_slots = 0 def available_size(self): - return len(self.free_slots) + return len(self.free_slots) + self.num_release_slots @synchronized def alloc(self, need_size: int) -> Optional[torch.Tensor]: @@ -260,13 +262,22 @@ class MambaPoolHost(HostKVCache): ), "The requested size should be a multiple of the page size." if need_size > self.available_size(): return None + + if need_size > len(self.free_slots): + self._merge_release_slots() + select_index = self.free_slots[:need_size] self.free_slots = self.free_slots[need_size:] return select_index @synchronized def free(self, indices: torch.Tensor) -> int: - self.free_slots = torch.cat([self.free_slots, indices]) + indices_cpu = indices.cpu() + if indices_cpu.numel() == 0: + return 0 + + self.release_slots.append(indices_cpu) + self.num_release_slots += len(indices_cpu) return len(indices) def get_size_per_token(self): @@ -850,9 +861,11 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache): def clear(self): self.free_slots = torch.arange(self.size, dtype=torch.int64) + self.release_slots = [] + self.num_release_slots = 0 def available_size(self): - return len(self.free_slots) + return len(self.free_slots) + self.num_release_slots @synchronized def alloc(self, need_size: int) -> Optional[torch.Tensor]: @@ -861,15 +874,22 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache): ) * self.slot_page_size if need_size > self.available_size(): return None + + if need_size > len(self.free_slots): + self._merge_release_slots() + select_index = self.free_slots[:need_size] self.free_slots = self.free_slots[need_size:] return select_index @synchronized def free(self, indices: torch.Tensor) -> int: - self.free_slots = torch.cat( - [self.free_slots, indices.to(dtype=torch.int64, device="cpu").flatten()] - ) + indices_cpu = indices.cpu() + if indices_cpu.numel() == 0: + return 0 + + self.release_slots.append(indices_cpu) + self.num_release_slots += len(indices_cpu) return len(indices) def backup_from_device_all_layer( diff --git a/python/sglang/srt/mem_cache/pool_host/base.py b/python/sglang/srt/mem_cache/pool_host/base.py index 5dba1aa78..84d5e6bd5 100644 --- a/python/sglang/srt/mem_cache/pool_host/base.py +++ b/python/sglang/srt/mem_cache/pool_host/base.py @@ -267,12 +267,28 @@ class HostKVCache(abc.ABC): (self.size,), dtype=torch.uint8, device=self.device ) self.free_slots = torch.arange(self.size, dtype=torch.int64) + # Keep freed chunks aside and consume them lazily from alloc() to avoid + # concatenating a large free-list on every host-pool free. + self.release_slots = [] + self.num_release_slots = 0 # Per-slot flag used to detect double-free. # slot_used[k] is true if slot k is allocated. self.slot_used = torch.zeros(self.size, dtype=torch.bool) def available_size(self): - return len(self.free_slots) + return len(self.free_slots) + self.num_release_slots + + def _merge_release_slots(self): + if self.num_release_slots == 0: + return + + if len(self.free_slots) == 0 and len(self.release_slots) == 1: + self.free_slots = self.release_slots[0] + else: + self.free_slots = torch.cat([self.free_slots, *self.release_slots]) + + self.release_slots = [] + self.num_release_slots = 0 @synchronized def alloc(self, need_size: int) -> Optional[torch.Tensor]: @@ -282,6 +298,9 @@ class HostKVCache(abc.ABC): if need_size > self.available_size(): return None + if need_size > len(self.free_slots): + self._merge_release_slots() + select_index = self.free_slots[:need_size] self.free_slots = self.free_slots[need_size:] @@ -296,10 +315,14 @@ class HostKVCache(abc.ABC): @synchronized def free(self, indices: torch.Tensor) -> int: indices_cpu = indices.cpu() + if indices_cpu.numel() == 0: + return 0 + assert self.slot_used[indices_cpu].all(), ( f"Double-free detected: slots not currently allocated: " f"{indices_cpu[~self.slot_used[indices_cpu]].tolist()}." ) self.slot_used[indices_cpu] = False - self.free_slots = torch.cat([self.free_slots, indices_cpu]) + self.release_slots.append(indices_cpu) + self.num_release_slots += len(indices_cpu) return len(indices) diff --git a/test/registered/unit/mem_cache/test_mem_pool_host.py b/test/registered/unit/mem_cache/test_mem_pool_host.py index da35d527c..510bf074c 100644 --- a/test/registered/unit/mem_cache/test_mem_pool_host.py +++ b/test/registered/unit/mem_cache/test_mem_pool_host.py @@ -1,10 +1,15 @@ -"""Unit tests for HostKVCache alloc/free bookkeeping (double-alloc / double-free detection).""" +"""Unit tests for host-pool allocation and free-list bookkeeping.""" +import threading import unittest import torch from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool +from sglang.srt.mem_cache.memory_pool_host import ( + DeepSeekV4PagedHostPool, + MambaPoolHost, +) from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -78,6 +83,86 @@ class TestHostKVCache(CustomTestCase): self.assertIn("Double-free", msg) self.assertIn(str(indices.tolist()), msg) + def test_empty_free_keeps_release_list_empty(self): + self.assertEqual(self.host_pool.free(torch.empty(0, dtype=torch.int64)), 0) + self.assertEqual(self.host_pool.num_release_slots, 0) + self.assertEqual(self.host_pool.release_slots, []) + + +class TestLazyHostPoolRelease(CustomTestCase): + @staticmethod + def _make_mamba_pool(): + pool = MambaPoolHost.__new__(MambaPoolHost) + pool.size = 8 + pool.page_size = 1 + pool.device = "cpu" + pool.lock = threading.RLock() + pool.clear() + return pool + + @staticmethod + def _make_deepseek_v4_pool(): + pool = DeepSeekV4PagedHostPool.__new__(DeepSeekV4PagedHostPool) + pool.size = 8 + pool.slot_page_size = 2 + pool.lock = threading.RLock() + pool.clear() + return pool + + def _assert_lazy_release(self, pool): + self.assertEqual(pool.free(torch.empty(0, dtype=torch.int64)), 0) + self.assertEqual(pool.num_release_slots, 0) + self.assertEqual(pool.release_slots, []) + + allocated = pool.alloc(6) + free_slots_before = pool.free_slots + + pool.free(allocated[:2]) + + # free() should keep the primary free-list untouched and only record + # the released chunk for a later merge. + self.assertIs(pool.free_slots, free_slots_before) + self.assertEqual(pool.num_release_slots, 2) + self.assertEqual(len(pool.release_slots), 1) + self.assertEqual(pool.available_size(), 4) + + # Consume the primary free-list first without merging pending slots. + self.assertTrue(torch.equal(pool.alloc(2), torch.tensor([6, 7]))) + self.assertEqual(pool.num_release_slots, 2) + + # Once the primary free-list is exhausted, alloc() merges and reuses + # the pending slots. + self.assertTrue(torch.equal(pool.alloc(2), torch.tensor([0, 1]))) + self.assertEqual(pool.num_release_slots, 0) + self.assertEqual(pool.release_slots, []) + self.assertEqual(pool.available_size(), 0) + + pool.free(torch.tensor([0, 1])) + pool.clear() + self.assertEqual(pool.num_release_slots, 0) + self.assertEqual(pool.release_slots, []) + self.assertEqual(pool.available_size(), 8) + + # Exercise the general merge path with multiple released chunks. + allocated = pool.alloc(8) + pool.free(allocated[:2]) + pool.free(allocated[2:4]) + self.assertEqual(len(pool.release_slots), 2) + self.assertTrue(torch.equal(pool.alloc(4), torch.tensor([0, 1, 2, 3]))) + self.assertEqual(pool.num_release_slots, 0) + self.assertEqual(pool.release_slots, []) + + def test_mamba_pool_lazy_release(self): + self._assert_lazy_release(self._make_mamba_pool()) + + def test_deepseek_v4_pool_lazy_release(self): + pool = self._make_deepseek_v4_pool() + self._assert_lazy_release(pool) + + # Preserve the pool's page-aligned allocation behavior. + pool.clear() + self.assertEqual(len(pool.alloc(1)), 2) + if __name__ == "__main__": unittest.main()