fix: preserve HiCache load-back allocator headroom

This commit is contained in:
2026-05-08 00:18:33 +08:00
parent 1f074f434e
commit 95bacb8862
10 changed files with 172 additions and 37 deletions

View File

@@ -499,6 +499,15 @@ class PrefillAdder:
def ceil_paged_tokens(self, tokens: int) -> int:
return -(-tokens // self.page_size) * self.page_size
def _get_available_device_tokens_for_load_back(self) -> int:
if self.is_hybrid_swa:
return self.token_to_kv_pool_allocator.full_available_size()
return self.token_to_kv_pool_allocator.available_size()
def _get_load_back_mem_quota(self, real_input_tokens: int) -> int:
reserve_tokens = real_input_tokens + self.page_size
return max(self._get_available_device_tokens_for_load_back() - reserve_tokens, 0)
def budget_state(self):
if self.rem_total_tokens <= 0 or self.cur_rem_tokens <= 0:
return AddReqResult.NO_TOKEN
@@ -777,6 +786,7 @@ class PrefillAdder:
InitLoadBackParams(
last_host_node=req.last_host_node,
host_hit_length=req.host_hit_length,
mem_quota=self._get_load_back_mem_quota(real_input_tokens),
)
)
req.prefix_indices = torch.cat([req.prefix_indices, new_indices])

View File

@@ -69,6 +69,38 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
def available_size(self):
return (len(self.free_pages) + len(self.release_pages)) * self.page_size
def immediate_available_pages(self) -> int:
if self.free_pages is None:
return 0
return len(self.free_pages)
def deferred_available_pages(self) -> int:
if self.release_pages is None:
return 0
return len(self.release_pages)
def immediate_available_size(self) -> int:
if self.free_pages is None:
return self.available_size()
return self.immediate_available_pages() * self.page_size
def deferred_available_size(self) -> int:
if self.release_pages is None:
return 0
return self.deferred_available_pages() * self.page_size
def allocator_state_str(self) -> str:
if self.free_pages is None:
return f"allocator_available_size={self.available_size()}"
return (
"allocator_available_size="
f"{self.available_size()} "
f"(allocator_free_size={self.immediate_available_size()} "
f"[free_pages={self.immediate_available_pages()}] + "
f"allocator_release_size={self.deferred_available_size()} "
f"[release_pages={self.deferred_available_pages()}])"
)
def get_kvcache(self):
return self._kvcache
@@ -434,10 +466,15 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
)
bs = len(prefix_lens)
if self.need_sort and extend_num_tokens // self.page_size + bs + 1 > len(
self.free_pages
):
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
prefix_lens=prefix_lens_cpu,
)
if self.need_sort and num_new_pages > len(self.free_pages):
self.merge_and_sort_free()
if num_new_pages > len(self.free_pages):
return None
out_indices = torch.empty(
(extend_num_tokens,), dtype=torch.int64, device=self.device
@@ -456,14 +493,6 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
if self.debug_mode:
assert len(torch.unique(out_indices)) == len(out_indices)
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
prefix_lens=prefix_lens_cpu,
)
if num_new_pages > len(self.free_pages):
return None
self.free_pages = self.free_pages[num_new_pages:]
return out_indices
@@ -479,8 +508,15 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
)
bs = len(seq_lens)
if self.need_sort and bs > len(self.free_pages):
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
decode=True,
)
if self.need_sort and num_new_pages > len(self.free_pages):
self.merge_and_sort_free()
if num_new_pages > len(self.free_pages):
return None
out_indices = torch.empty((bs,), dtype=torch.int64, device=self.device)
alloc_decode_kernel[(bs,)](
@@ -495,14 +531,6 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
if self.debug_mode:
assert len(torch.unique(out_indices)) == len(out_indices)
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
decode=True,
)
if num_new_pages > len(self.free_pages):
return None
self.free_pages = self.free_pages[num_new_pages:]
return out_indices

View File

@@ -268,6 +268,14 @@ class BasePrefixCache(ABC, PrefixCacheTrait):
return not self.is_chunk_cache()
def available_and_evictable_str(self) -> str:
available_size = self.token_to_kv_pool_allocator.available_size()
allocator = self.token_to_kv_pool_allocator
allocator_state_str = getattr(allocator, "allocator_state_str", None)
if allocator_state_str is None:
allocator_state = f"allocator_available_size={allocator.available_size()}"
else:
allocator_state = allocator_state_str()
evictable_size = self.evictable_size()
return f"Available tokens: {available_size + evictable_size} ({available_size=} + {evictable_size=})\n"
return (
f"Available tokens: {allocator.available_size() + evictable_size} "
f"({allocator_state} + {evictable_size=})\n"
)

View File

@@ -7,7 +7,11 @@ import torch
import triton
import triton.language as tl
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache,
EvictParams,
EvictResult,
)
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_in_seq_page_compute_owners,
get_in_seq_page_compute_owner_unavailable_reason,
@@ -220,7 +224,7 @@ def alloc_token_slots(
backup_state: bool = False,
):
allocator = tree_cache.token_to_kv_pool_allocator
evict_from_tree_cache(tree_cache, num_tokens)
evict_result = evict_from_tree_cache(tree_cache, num_tokens)
state = None
if backup_state:
@@ -233,6 +237,7 @@ def alloc_token_slots(
f"Out of memory. Try to lower your batch size.\n"
f"Try to allocate {num_tokens} tokens.\n"
f"{available_and_evictable_str(tree_cache)}"
f"{_evict_result_str(evict_result)}\n"
)
logger.error(error_msg)
if tree_cache is not None:
@@ -242,12 +247,25 @@ def alloc_token_slots(
return (out_cache_loc, state) if backup_state else out_cache_loc
def evict_from_tree_cache(tree_cache: BasePrefixCache | None, num_tokens: int):
def _evict_result_str(evict_result: EvictResult | None) -> str:
if evict_result is None:
return "evict_result=None"
return (
"evict_result=("
f"num_tokens_evicted={evict_result.num_tokens_evicted}, "
f"swa_num_tokens_evicted={evict_result.swa_num_tokens_evicted}, "
f"mamba_num_evicted={evict_result.mamba_num_evicted})"
)
def evict_from_tree_cache(
tree_cache: BasePrefixCache | None, num_tokens: int
) -> EvictResult:
if tree_cache is None:
return
return EvictResult()
if tree_cache.is_chunk_cache():
return
return EvictResult()
allocator = tree_cache.token_to_kv_pool_allocator
@@ -259,13 +277,15 @@ def evict_from_tree_cache(tree_cache: BasePrefixCache | None, num_tokens: int):
if full_available_size < num_tokens or swa_available_size < num_tokens:
full_num_tokens = max(0, num_tokens - full_available_size)
swa_num_tokens = max(0, num_tokens - swa_available_size)
tree_cache.evict(
return tree_cache.evict(
EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
)
return EvictResult()
else:
# Standard allocator
if allocator.available_size() < num_tokens:
tree_cache.evict(EvictParams(num_tokens=num_tokens))
return tree_cache.evict(EvictParams(num_tokens=num_tokens))
return EvictResult()
def _evict_for_compute_owner_lanes(
@@ -322,7 +342,7 @@ def alloc_paged_token_slots_extend(
# Over estimate the number of tokens: assume each request needs a new page.
allocator = tree_cache.token_to_kv_pool_allocator
num_tokens = extend_num_tokens + len(seq_lens_cpu) * allocator.page_size
evict_from_tree_cache(tree_cache, num_tokens)
evict_result = evict_from_tree_cache(tree_cache, num_tokens)
alloc_extend_compute_owner = getattr(
allocator, "alloc_extend_compute_owner", None
@@ -448,6 +468,7 @@ def alloc_paged_token_slots_extend(
f"Prefill out of memory. Try to lower your batch size.\n"
f"Try to allocate {extend_num_tokens} tokens.\n"
f"{available_and_evictable_str(tree_cache)}"
f"{_evict_result_str(evict_result)}\n"
)
logger.error(error_msg)
if tree_cache is not None:

View File

@@ -20,7 +20,6 @@ The radix tree data structure for managing the hybrid (full and Mamba) KV cache.
"""
import heapq
from collections import defaultdict
from functools import lru_cache, partial
from typing import TYPE_CHECKING, List, Optional, Tuple
@@ -69,7 +68,7 @@ class TreeNode:
last_access_time_counter_float = float64(1.0)
def __init__(self, id: Optional[int] = None):
self.children = defaultdict(TreeNode)
self.children = {}
self.parent: TreeNode = None
self.key: RadixKey = None
self.value: Optional[torch.Tensor] = None

View File

@@ -26,7 +26,6 @@ import heapq
import logging
import sys
import time
from collections import defaultdict
from functools import lru_cache, partial
from typing import TYPE_CHECKING, Any, Iterator, List, Optional, Tuple, Union
@@ -123,7 +122,7 @@ class TreeNode:
counter = 0
def __init__(self, id: Optional[int] = None, priority: int = 0):
self.children = defaultdict(TreeNode)
self.children = {}
self.parent: TreeNode = None
self.key: RadixKey = None
self.value: Optional[torch.Tensor] = None

View File

@@ -21,7 +21,6 @@ The radix tree data structure for managing the hybrid (full and SWA) KV cache.
import heapq
import time
from collections import defaultdict
from functools import partial
from typing import TYPE_CHECKING, List, Optional, Tuple
@@ -67,7 +66,7 @@ class TreeNode:
last_access_time_counter_float = float64(1.0)
def __init__(self, id: Optional[int] = None):
self.children = defaultdict(TreeNode)
self.children = {}
self.parent: TreeNode = None
self.key: RadixKey = None
self.value: Optional[torch.Tensor] = None

View File

@@ -442,6 +442,40 @@ class TestPrefillAdder(CustomTestCase):
self.assertEqual(adder2.rem_chunk_tokens, 0) # 3 - 3 = 0
self.assertEqual(result3, AddReqResult.OTHER)
def test_host_load_back_passes_mem_quota(self):
running_batch = self.create_running_batch()
self.mock_token_allocator.available_size.return_value = 512
self.mock_tree_cache.init_load_back.return_value = (
__import__("torch").tensor([1, 2, 3, 4], dtype=__import__("torch").int64),
"loaded_node",
)
adder = self.create_adder(
running_batch,
page_size=64,
rem_input_tokens=4096,
rem_total_tokens=4096,
)
req = self.create_mock_req("req", priority=0, max_new_tokens=16)
req.extend_input_len = 256
req.host_hit_length = 128
req.prefix_indices = __import__("torch").empty(
(0,), dtype=__import__("torch").int64
)
req.last_node = object()
req.last_host_node = object()
req.fill_ids = list(range(256))
req.cache_protected_len = 0
req.set_extend_input_len = lambda value: setattr(req, "extend_input_len", value)
req.sampling_params.ignore_eos = False
result = adder.add_one_req(
req, has_chunked_req=False, truncation_align_size=None
)
self.assertNotEqual(result, AddReqResult.NO_TOKEN)
params = self.mock_tree_cache.init_load_back.call_args.args[0]
self.assertEqual(params.mem_quota, 256)
if __name__ == "__main__":
unittest.main()

View File

@@ -6,6 +6,7 @@ import torch
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
@@ -52,7 +53,7 @@ class TestCpSharedKVLayout(unittest.TestCase):
self.assertEqual(positions.tolist(), [11, 15])
class TestCPSharedPagedAllocator(unittest.TestCase):
class TestCPSharedPagedAllocator(CustomTestCase):
def test_compute_owner_alloc_fallback_logs_every_event(self):
from sglang.srt.mem_cache import common
@@ -95,6 +96,31 @@ class TestCPSharedPagedAllocator(unittest.TestCase):
allocator.free(locs)
self.assertEqual(allocator.available_size(), 64 * 8)
def test_allocator_state_str_reports_free_and_release_pages(self):
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
allocator = CPSharedPagedTokenToKVPoolAllocator(
logical_size=64 * 8,
physical_size=64 * 2,
page_size=64,
dtype=torch.bfloat16,
device="cpu",
kvcache=None,
need_sort=True,
cp_size=4,
cp_rank=0,
)
allocator.free_pages = torch.tensor([1, 5], dtype=torch.int64)
allocator.release_pages = torch.tensor([9], dtype=torch.int64)
state = allocator.allocator_state_str()
self.assertIn("allocator_available_size=192", state)
self.assertIn("allocator_free_size=128", state)
self.assertIn("free_pages=2", state)
self.assertIn("allocator_release_size=64", state)
self.assertIn("release_pages=1", state)
def test_compute_owner_page_assignment_matches_in_seq_zigzag(self):
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_in_seq_page_compute_owners,

View File

@@ -39,6 +39,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams,
)
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
from sglang.test.test_utils import CustomTestCase
# Test constants
DEFAULT_PAGE_SIZE = 4
@@ -235,6 +236,16 @@ class TestTreeNode(unittest.TestCase):
self.assertFalse(node2 < node1)
class TestTreeNodeChildren(CustomTestCase):
def test_missing_child_access_does_not_create_orphan_node(self):
node = TreeNode()
with self.assertRaises(KeyError):
_ = node.children["missing"]
self.assertEqual(node.children, {})
class TestRadixCache(unittest.TestCase):
"""Test cases for RadixCache class."""