Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
95 changes: 95 additions & 0 deletions kvcached/control.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,95 @@
# SPDX-FileCopyrightText: Copyright contributors to the kvcached project
# SPDX-License-Identifier: Apache-2.0

"""Revisioned memory controls for live kvcached KV pools."""

from __future__ import annotations

from typing import Any, Dict

from kvcached.pool_registry import get_registered_kv_cache_pools


def set_instance_memory_limit(
limit_bytes: int,
*,
revision: int,
) -> Dict[str, Any]:
"""Split and apply one instance limit across all live KV pools."""
limit_bytes = int(limit_bytes)
revision = int(revision)
if limit_bytes < 0:
raise ValueError("limit_bytes must be non-negative")
if revision < 0:
raise ValueError("revision must be non-negative")

managers = [manager for manager, _ in get_registered_kv_cache_pools()]
if not managers:
return {
"status": "unavailable",
"reason": "no_registered_kv_cache_pool",
"limit_bytes": limit_bytes,
"effective_limit_bytes": 0,
"revision": revision,
"mapped_bytes": 0,
"remaining_bytes": 0,
"overage_bytes": 0,
"pools": [],
}

managers.sort(
key=lambda manager: (
str(manager.pool_name or ""),
int(manager.group_id),
)
)
capacities = [
max(
0,
int(manager.mem_size) * int(manager.num_layers) * int(manager.num_kv_buffers),
)
for manager in managers
]
total_capacity = sum(capacities)
if total_capacity <= 0:
raise ValueError("registered KV pools have no virtual capacity")

remaining = min(limit_bytes, total_capacity)
pool_states = []
for index, (manager, capacity) in enumerate(zip(managers, capacities)):
share = (
remaining
if index == len(managers) - 1
else min(
remaining,
min(limit_bytes, total_capacity) * capacity // total_capacity,
)
)
pool_states.append(manager.set_memory_limit(share, revision=revision))
remaining -= share

mapped = sum(int(state["mapped_bytes"]) for state in pool_states)
effective = sum(int(state["effective_limit_bytes"] or 0) for state in pool_states)
statuses = {str(state["status"]) for state in pool_states}
if "conflict" in statuses:
status = "conflict"
elif "stale" in statuses:
status = "stale"
elif "deferred" in statuses:
status = "deferred"
else:
status = "applied"
return {
"status": status,
"reason": {
"deferred": "inuse_capacity_above_limit",
"conflict": "revision_reused_with_different_limit",
}.get(status, ""),
"limit_bytes": limit_bytes,
"effective_limit_bytes": effective,
"revision": revision,
"mapped_bytes": mapped,
"remaining_bytes": max(0, effective - mapped),
"overage_bytes": max(0, mapped - effective),
"pools": pool_states,
}
86 changes: 85 additions & 1 deletion kvcached/kv_cache_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,6 +186,9 @@ def unmap_callback(world_size: int, offsets: List[int]) -> None:

self.in_shrink: bool = False
self.target_num_blocks: Optional[int] = None
self._memory_limit_bytes: Optional[int] = None
self._memory_limit_effective_bytes: Optional[int] = None
self._memory_limit_revision = -1
# NOTE: we use a no-op lock for sync scheduling to avoid overhead
self._lock = threading.RLock() if async_sched else NoOpLock()

Expand Down Expand Up @@ -467,7 +470,7 @@ def resize(self, new_mem_size: int):
new_mem_size: the memory size of the K or V tensor in one layer
"""
self._wait_post_init()
assert new_mem_size > 0, "new_mem_size must be positive"
assert new_mem_size >= 0, "new_mem_size must be non-negative"
if self.page_allocator.resize(new_mem_size):
if self.in_shrink:
self.in_shrink = False
Expand All @@ -492,6 +495,87 @@ def trim(self) -> None:
self._wait_post_init()
self.page_allocator.trim()

@synchronized
def set_memory_limit(
self,
limit_bytes: int,
*,
revision: int,
) -> Dict[str, Any]:
"""Apply a revisioned memory limit through the existing resize path."""
limit_bytes = int(limit_bytes)
revision = int(revision)
if limit_bytes < 0:
raise ValueError("limit_bytes must be non-negative")
if revision < 0:
raise ValueError("revision must be non-negative")

current_revision = self._memory_limit_revision
current_limit = self._memory_limit_bytes
if revision < current_revision:
return self._memory_limit_state(status="stale")
if revision == current_revision and current_limit == limit_bytes:
return self._memory_limit_state()
if revision == current_revision:
return self._memory_limit_state(status="conflict")

page_bundle_bytes = self._memory_limit_page_bundle_bytes()
max_pages = self.mem_size // self.page_size
target_pages = min(limit_bytes // page_bundle_bytes, max_pages)
effective_limit_bytes = target_pages * page_bundle_bytes

self.resize(target_pages * self.page_size)
self._memory_limit_bytes = limit_bytes
self._memory_limit_effective_bytes = effective_limit_bytes
self._memory_limit_revision = revision
return self._memory_limit_state()

def _memory_limit_page_bundle_bytes(self) -> int:
return self.page_size * self.num_layers * self.num_kv_buffers

def _memory_limit_state(
self,
*,
status: Optional[str] = None,
) -> Dict[str, Any]:
page_state = self.page_allocator.get_page_state()
page_bundle_bytes = self._memory_limit_page_bundle_bytes()
mapped_pages = (int(page_state["inuse_pages"])
+ int(page_state["reserved_pages"]))
mapped_bytes = mapped_pages * page_bundle_bytes
effective_limit_bytes = self._memory_limit_effective_bytes
if status is None:
status = "deferred" if self.in_shrink else "applied"
return {
"status": status,
"pool_name": str(self.pool_name or ""),
"group_id": self.group_id,
"limit_bytes": self._memory_limit_bytes,
"effective_limit_bytes": effective_limit_bytes,
"current_capacity_bytes": (
int(page_state["total_pages"]) * page_bundle_bytes
),
"revision": self._memory_limit_revision,
"mapped_bytes": mapped_bytes,
"remaining_bytes": (
None if effective_limit_bytes is None else
max(0, effective_limit_bytes - mapped_bytes)
),
"overage_bytes": (
0 if effective_limit_bytes is None else
max(0, mapped_bytes - effective_limit_bytes)
),
"reason": {
"deferred": "inuse_capacity_above_limit",
"conflict": "revision_reused_with_different_limit",
}.get(status, ""),
}

@synchronized
def memory_limit_state(self) -> Dict[str, Any]:
"""Return the current revisioned resize limit and apply state."""
return self._memory_limit_state()

@synchronized
def available_size(self) -> int:
avail_blocks = self.num_avail_blocks + len(self.reserved_blocks)
Expand Down
1 change: 1 addition & 0 deletions tests/manifests/cpu.txt
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ tests/test_ipc_name.py
tests/test_ipc_timeout.py
tests/test_kv_cache_shape_compat.py
tests/test_make_cache_key.py
tests/test_memory_limit_control.py
tests/test_observability.py
tests/test_page_aware_eviction.py
tests/test_prefix_cache.py
Expand Down
Loading
Loading