From 7a6c652c778723378556de1003beec2fd92a131b Mon Sep 17 00:00:00 2001 From: Zhiqiang Xie Date: Sat, 19 Sep 2026 12:50:43 -0700 Subject: [PATCH] [HiCache] Auto-size the host pool to fit available host memory (#40135) --- python/sglang/srt/arg_groups/fields/memory.py | 4 + python/sglang/srt/arg_groups/hicache_hook.py | 19 ++- .../sglang/srt/mem_cache/hicache_auto_size.py | 123 +++++++++++++++++ python/sglang/srt/mem_cache/host_memory.py | 104 ++++++++++++++ .../sglang/srt/mem_cache/kv_cache_builder.py | 53 +++---- python/sglang/srt/mem_cache/pool_host/base.py | 40 +++++- python/sglang/srt/mem_cache/pool_host/dsa.py | 2 +- .../sglang/srt/mem_cache/pool_host/mamba.py | 2 +- python/sglang/srt/mem_cache/pool_host/mha.py | 2 +- .../entrypoints/openai/test_serving_chat.py | 20 +++ .../unit/mem_cache/test_hicache_auto_size.py | 104 ++++++++++++++ .../unit/mem_cache/test_host_memory.py | 130 ++++++++++++++++++ .../unit/mem_cache/test_mem_pool_host.py | 3 +- 13 files changed, 564 insertions(+), 42 deletions(-) create mode 100644 python/sglang/srt/mem_cache/hicache_auto_size.py create mode 100644 python/sglang/srt/mem_cache/host_memory.py create mode 100644 test/registered/unit/mem_cache/test_hicache_auto_size.py create mode 100644 test/registered/unit/mem_cache/test_host_memory.py diff --git a/python/sglang/srt/arg_groups/fields/memory.py b/python/sglang/srt/arg_groups/fields/memory.py index a3a046197..29c2b49dd 100644 --- a/python/sglang/srt/arg_groups/fields/memory.py +++ b/python/sglang/srt/arg_groups/fields/memory.py @@ -119,6 +119,10 @@ class Memory(msgspec.Struct): int, "The size of host KV cache memory pool in gigabytes. Overrides --hicache-ratio in either host memory mode.", ] = 0 + hicache_host_memory_fraction: A[ + Optional[float], + "Fraction of the available host memory, bounded by visible cgroup memory.max/memory.high or v1 memory limits (after a 10 GiB reserve) that the HiCache host pools of all ranks on this machine may use. Applies only when neither --hicache-ratio nor --hicache-size is set: the default ratio is then reduced until the pools fit. Lower it when several engines share a memory cgroup.", + ] = 0.8 hicache_write_policy: A[ str, Arg( diff --git a/python/sglang/srt/arg_groups/hicache_hook.py b/python/sglang/srt/arg_groups/hicache_hook.py index 57ed5eb7b..f9e4df805 100644 --- a/python/sglang/srt/arg_groups/hicache_hook.py +++ b/python/sglang/srt/arg_groups/hicache_hook.py @@ -69,16 +69,23 @@ def handle_hicache_ratio_default(server_args: Any): A decode server keeps the ratio unset here: kv_cache_builder resolves it against the retraction-backup backend (1.0 for host_pool, else 2.0). + + An explicit --hicache-ratio or --hicache-size is honored as given, so it + resolves --hicache-host-memory-fraction to None (auto-sizing off). """ cfg = resolving_view(server_args) + fraction = cfg.hicache_host_memory_fraction + if fraction is not None and not 0 < fraction <= 1: + raise ValueError("--hicache-host-memory-fraction must be in (0, 1].") + fields = {} if cfg.hicache_ratio is None and cfg.disaggregation_mode != "decode": - declare_resolution( - server_args, - "_handle_hicache_ratio_default", - hicache_ratio=( - 1.2 if cfg.hicache_host_memory_mode == "buffer_only" else 2.0 - ), + fields["hicache_ratio"] = ( + 1.2 if cfg.hicache_host_memory_mode == "buffer_only" else 2.0 ) + if cfg.hicache_ratio is not None or cfg.hicache_size > 0: + fields["hicache_host_memory_fraction"] = None + if fields: + declare_resolution(server_args, "_handle_hicache_ratio_default", **fields) def resolve_hicache_dcp_compatibility(server_args: Any): diff --git a/python/sglang/srt/mem_cache/hicache_auto_size.py b/python/sglang/srt/mem_cache/hicache_auto_size.py new file mode 100644 index 000000000..91d613c25 --- /dev/null +++ b/python/sglang/srt/mem_cache/hicache_auto_size.py @@ -0,0 +1,123 @@ +from __future__ import annotations + +import logging +from contextlib import contextmanager +from typing import TYPE_CHECKING + +import torch + +from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool +from sglang.srt.mem_cache.memory_pool import ( + HybridLinearKVPool, + HybridReqToTokenPool, + MHATokenToKVPool, + MiniMaxSparseKVPool, + MLATokenToKVPool, +) +from sglang.srt.mem_cache.pool_host.base import ( + host_memory_budget_bytes, + host_memory_budget_scope, + ranks_per_host, +) +from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool +from sglang.srt.runtime_context import get_context, get_memory, get_parallel + +if TYPE_CHECKING: + from sglang.srt.mem_cache.cache_init_params import CacheInitParams + from sglang.srt.speculative.base_spec_worker import HiCacheDraftPlan + +logger = logging.getLogger(__name__) + +# Page rounding, allocator metadata and staging buffers are outside the device +# byte counts the ratio is derived from. +_ALLOCATION_SLACK_FRACTION = 0.05 + +_SIZEABLE_POOLS = ( + MHATokenToKVPool, + MLATokenToKVPool, + SWAKVPool, + HybridLinearKVPool, + MiniMaxSparseKVPool, +) + + +def _pool_bytes(pool) -> int: + if isinstance(pool, SWAKVPool): + return _pool_bytes(pool.full_kv_pool) + _pool_bytes(pool.swa_kv_pool) + if isinstance(pool, HybridLinearKVPool): + return _pool_bytes(pool.full_kv_pool) + sizes = pool.get_kv_size_bytes() + return sum(sizes) if isinstance(sizes, tuple) else sizes + + +def _draft_bytes(target, draft) -> int: + if isinstance(draft, BaseSWAKVPool): + # Match sidecar construction: only SWA drafts follow target SWA slots. + target, draft = target.swa_kv_pool, draft.swa_kv_pool + # A sidecar has one host slot per target slot, however few slots the draft has. + return _pool_bytes(draft) * target.size // draft.size + + +def _estimate_hicache_bytes( + params: CacheInitParams, draft_plan: HiCacheDraftPlan | None +) -> int: + """Device bytes whose host mirrors scale with the HiCache ratio.""" + pool = params.token_to_kv_pool_allocator.get_kvcache() + if not isinstance(pool, _SIZEABLE_POOLS): + raise ValueError( + f"HiCache auto-sizing does not support {type(pool).__name__}; " + "set --hicache-ratio or --hicache-size explicitly." + ) + total = _pool_bytes(pool) + if isinstance(params.req_to_token_pool, HybridReqToTokenPool): + total += _pool_bytes(params.req_to_token_pool.mamba_pool) + drafts = params.mtp_draft_device_pools + if draft_plan is not None and draft_plan.mode == "sidecar": + drafts = draft_plan.device_pools + return total + sum(_draft_bytes(pool, draft) for draft in drafts) + + +@contextmanager +def auto_size_hicache( + params: CacheInitParams, draft_plan: HiCacheDraftPlan | None, *, enabled: bool +): + """Reduce the default HiCache ratio until this machine's host pools fit. + + Resolution nulls the fraction for an explicit --hicache-ratio/--hicache-size. + """ + fraction = get_memory().hicache_host_memory_fraction + if not enabled or fraction is None: + yield + return + requested = get_memory().hicache_ratio + device_bytes = _estimate_hicache_bytes(params, draft_plan) + budget = int(host_memory_budget_bytes() * fraction) + ratio = min(requested, budget * (1 - _ALLOCATION_SLACK_FRACTION) / device_bytes) + # One collective before any pool is built: PP stages own different pool + # counts, so a per-pool collective could deadlock. + if torch.distributed.is_initialized(): + value = torch.tensor([ratio], dtype=torch.float64) + torch.distributed.all_reduce( + value, + op=torch.distributed.ReduceOp.MIN, + group=get_parallel().world_group.cpu_group, + ) + ratio = value.item() + if ratio <= 0: + raise ValueError( + "No host memory is left for HiCache after the 10 GiB reserve; " + "set --hicache-ratio or --hicache-size explicitly." + ) + get_context().override("hicache.auto_size", hicache_ratio=ratio) + logger.info( + "HiCache auto-sizing: ratio %.3f -> %.3f; %.1f GiB host memory per rank " + "(fraction %.2f, %d ranks on this host), host pools %.1f GiB.", + requested, + ratio, + budget / 1024**3, + fraction, + ranks_per_host(), + device_bytes * ratio / 1024**3, + ) + with host_memory_budget_scope(budget): + yield diff --git a/python/sglang/srt/mem_cache/host_memory.py b/python/sglang/srt/mem_cache/host_memory.py new file mode 100644 index 000000000..8cd8b4132 --- /dev/null +++ b/python/sglang/srt/mem_cache/host_memory.py @@ -0,0 +1,104 @@ +"""Host-memory headroom bounded by the process's visible cgroup hierarchy.""" + +from __future__ import annotations + +import logging +import re +from pathlib import Path, PurePosixPath + +import psutil + +logger = logging.getLogger(__name__) + + +def _unescape_mount_path(value: str) -> str: + return re.sub(r"\\([0-7]{3})", lambda m: chr(int(m[1], 8)), value) + + +def _cgroup_memory_headroom(proc_root: Path = Path("/proc")) -> int | None: + memberships = {} + try: + cgroups = (proc_root / "self/cgroup").read_text() + mounts = (proc_root / "self/mountinfo").read_text() + except FileNotFoundError: + # Non-Linux systems need not expose procfs. + return None + for line in cgroups.splitlines(): + _, controllers, path = line.split(":", 2) + if not controllers: + memberships["cgroup2"] = PurePosixPath(path) + elif "memory" in controllers.split(","): + memberships["cgroup"] = PurePosixPath(path) + + headroom = None + resolved = False + for line in mounts.splitlines(): + before, after = line.split(" - ", 1) + filesystem, _, options = after.split()[:3] + if filesystem not in memberships: + continue + if filesystem == "cgroup" and "memory" not in options.split(","): + continue + fields = before.split() + root = PurePosixPath(_unescape_mount_path(fields[3])) + mount = Path(_unescape_mount_path(fields[4])) + membership = memberships[filesystem] + if membership.is_relative_to(root): + relative = membership.relative_to(root) + elif root != PurePosixPath("/"): + # A cgroup namespace can expose membership relative to its root, + # while mountinfo still identifies the host-side subtree. + relative = membership.relative_to("/") + else: + continue + if ".." in relative.parts: + raise ValueError(f"Cannot resolve cgroup memory path: {membership}") + directory = mount / relative + if not directory.is_dir(): + continue + resolved = True + limits = ( + ("memory.max", "memory.high") + if filesystem == "cgroup2" + else ("memory.limit_in_bytes",) + ) + usage_name = ( + "memory.current" if filesystem == "cgroup2" else "memory.usage_in_bytes" + ) + while True: + for name in limits: + try: + value = (directory / name).read_text().strip() + except FileNotFoundError: + # The hierarchy root may not have memory controller files. + continue + if value == "max": + continue + limit = int(value) + # Do not silently ignore an unreadable usage file for a known + # limit: falling back to host RAM could overrun the container. + usage = int((directory / usage_name).read_text()) + remaining = max(0, limit - usage) + headroom = remaining if headroom is None else min(headroom, remaining) + if directory == mount: + break + directory = directory.parent + if memberships and not resolved: + raise RuntimeError( + "Cannot locate the process memory cgroup in mounted cgroup filesystems" + ) + return headroom + + +def available_host_memory_bytes() -> int: + """Conservative allocatable RAM; charged file cache is not assumed reclaimable.""" + available = psutil.virtual_memory().available + cgroup_headroom = _cgroup_memory_headroom() + if cgroup_headroom is not None: + logger.info( + "HiCache memory headroom: host %.1f GiB, cgroup %.1f GiB", + available / 1024**3, + cgroup_headroom / 1024**3, + ) + available = min(available, cgroup_headroom) + return available diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index bee69557c..0007d143f 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -40,6 +40,7 @@ from sglang.srt.hardware_backend.mlx.runtime import use_mlx from sglang.srt.managers.mm_schedule import init_mm_embedding_cache from sglang.srt.mem_cache.cache_init_params import CacheInitParams from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool +from sglang.srt.mem_cache.hicache_auto_size import auto_size_hicache from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool from sglang.srt.mem_cache.registry import TreeCacheBuildContext, create_tree_cache from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool @@ -327,32 +328,36 @@ def build_kv_cache( mtp_draft_device_pools=mtp_draft_device_pools, ) - tree_cache = create_tree_cache( - TreeCacheBuildContext( - server_args=server_args, - params=params, - is_hybrid_swa=is_hybrid_swa, - full_tokens_per_layer=full_tokens_per_layer, - is_hybrid_ssm=is_hybrid_ssm, - is_dsa=is_dsa, - enable_hierarchical_cache=enable_hierarchical_cache, - disable_radix_cache=disable_radix_cache, - effective_chunked_prefill_size=effective_chunked_prefill_size, - tp_worker=tp_worker, - model_config=model_config, - tp_size=ps.tp_size, - tp_rank=ps.tp_rank, - tp_group=tp_group, - ) + tree_context = TreeCacheBuildContext( + server_args=server_args, + params=params, + is_hybrid_swa=is_hybrid_swa, + full_tokens_per_layer=full_tokens_per_layer, + is_hybrid_ssm=is_hybrid_ssm, + is_dsa=is_dsa, + enable_hierarchical_cache=enable_hierarchical_cache, + disable_radix_cache=disable_radix_cache, + effective_chunked_prefill_size=effective_chunked_prefill_size, + tp_worker=tp_worker, + model_config=model_config, + tp_size=ps.tp_size, + tp_rank=ps.tp_rank, + tp_group=tp_group, ) + with auto_size_hicache( + params, + hicache_draft_plan, + enabled=enable_hierarchical_cache or retraction_backup == "host_pool", + ): + tree_cache = create_tree_cache(tree_context) - if ( - enable_hierarchical_cache or retraction_backup == "host_pool" - ) and hicache_draft_plan is not None: - maybe_register_hicache_draft( - tree_cache=tree_cache, - draft_plan=hicache_draft_plan, - ) + if ( + enable_hierarchical_cache or retraction_backup == "host_pool" + ) and hicache_draft_plan is not None: + maybe_register_hicache_draft( + tree_cache=tree_cache, + draft_plan=hicache_draft_plan, + ) if retraction_backup == "host_pool": if not isinstance(tree_cache, UnifiedRadixCache): diff --git a/python/sglang/srt/mem_cache/pool_host/base.py b/python/sglang/srt/mem_cache/pool_host/base.py index 6a0261623..6d3418193 100644 --- a/python/sglang/srt/mem_cache/pool_host/base.py +++ b/python/sglang/srt/mem_cache/pool_host/base.py @@ -3,12 +3,14 @@ from __future__ import annotations import abc import logging import threading +from contextlib import contextmanager +from contextvars import ContextVar from functools import wraps from typing import Optional -import psutil import torch +from sglang.srt.mem_cache.host_memory import available_host_memory_bytes from sglang.srt.mem_cache.memory_pool import KVCache from sglang.srt.mem_cache.pool_host.common import ( _cuda_host_unregister, @@ -28,6 +30,21 @@ HICACHE_HOST_MEMORY_RESERVE_BYTES: int = 10 * (1024**3) _WRITE_BACK_STAGING_PAGE_CHUNK = 64 +_host_memory_budget: ContextVar[Optional[int]] = ContextVar( + "hicache_host_memory_budget", default=None +) + + +@contextmanager +def host_memory_budget_scope(budget_bytes: int): + """Book every pool built inside against one snapshot, not re-sampled psutil.""" + token = _host_memory_budget.set(budget_bytes) + try: + yield + finally: + _host_memory_budget.reset(token) + + def ranks_per_host() -> int: """Number of ranks of this job running on the same machine as this one. @@ -48,14 +65,23 @@ def ranks_per_host() -> int: return max(launch_world_size // get_parallel().nnodes, 1) -def host_memory_budget_bytes() -> int: +def host_memory_budget_bytes(requested_bytes: int = 0) -> int: """Host RAM this rank may claim for a HiCache pool. - psutil reports the whole machine, so co-located ranks each see the same free - memory; without the split every rank sizes its pool against all of it and - the host is oversubscribed by the number of ranks it holds. + Bound machine availability by the visible cgroup limits before splitting + among local ranks. Independent engines with separate container budgets + therefore size against their own remaining allowance. + + Inside host_memory_budget_scope, requested_bytes is booked against the + snapshot when it fits; the allowance before booking is returned. """ - free = psutil.virtual_memory().available - HICACHE_HOST_MEMORY_RESERVE_BYTES + available = _host_memory_budget.get() + if available is not None: + if requested_bytes <= available: + _host_memory_budget.set(available - requested_bytes) + return available + + free = available_host_memory_bytes() - HICACHE_HOST_MEMORY_RESERVE_BYTES return free // ranks_per_host() @@ -172,7 +198,7 @@ class HostKVCache(abc.ABC): # Verify there is enough available host memory. requested_bytes = self.size * self.size_per_token - available_bytes = host_memory_budget_bytes() + available_bytes = host_memory_budget_bytes(requested_bytes) if requested_bytes > available_bytes: raise ValueError( f"Not enough host memory available. Requesting " diff --git a/python/sglang/srt/mem_cache/pool_host/dsa.py b/python/sglang/srt/mem_cache/pool_host/dsa.py index 4e6ab2f16..e6d3d6241 100644 --- a/python/sglang/srt/mem_cache/pool_host/dsa.py +++ b/python/sglang/srt/mem_cache/pool_host/dsa.py @@ -111,7 +111,7 @@ class DSAIndexerPoolHost(HostKVCache): buf_elem_size = self.page_num * self.layer_num * self.indexer_page_stride_size requested_bytes = buf_elem_size * self.indexer_dtype.itemsize - available_bytes = host_memory_budget_bytes() + available_bytes = host_memory_budget_bytes(requested_bytes) if requested_bytes > available_bytes: raise ValueError( f"Not enough host memory for DSA indexer hierarchical cache. " diff --git a/python/sglang/srt/mem_cache/pool_host/mamba.py b/python/sglang/srt/mem_cache/pool_host/mamba.py index e8656ec15..328443837 100644 --- a/python/sglang/srt/mem_cache/pool_host/mamba.py +++ b/python/sglang/srt/mem_cache/pool_host/mamba.py @@ -129,7 +129,7 @@ class MambaPoolHost(HostKVCache): ) requested_bytes = self.size * self.size_per_token - available_bytes = host_memory_budget_bytes() + available_bytes = host_memory_budget_bytes(requested_bytes) if requested_bytes > available_bytes: raise ValueError( f"Not enough host memory available. Requesting " diff --git a/python/sglang/srt/mem_cache/pool_host/mha.py b/python/sglang/srt/mem_cache/pool_host/mha.py index 03e1629b1..f55904b61 100644 --- a/python/sglang/srt/mem_cache/pool_host/mha.py +++ b/python/sglang/srt/mem_cache/pool_host/mha.py @@ -765,7 +765,7 @@ class MHATokenToKOnlyPoolHost(HostKVCache): self.size_per_token = self.get_size_per_token() requested_bytes = self.size * self.size_per_token - available_bytes = host_memory_budget_bytes() + available_bytes = host_memory_budget_bytes(requested_bytes) if requested_bytes > available_bytes: raise ValueError( f"Not enough host memory for MiniMax index-K hierarchical cache. " diff --git a/test/registered/unit/entrypoints/openai/test_serving_chat.py b/test/registered/unit/entrypoints/openai/test_serving_chat.py index 18a87f036..343f16354 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_chat.py +++ b/test/registered/unit/entrypoints/openai/test_serving_chat.py @@ -10,6 +10,8 @@ from sglang.test.test_utils import CustomTestCase, enter_override, maybe_stub_sg maybe_stub_sgl_kernel() # must precede any import that pulls in sgl_kernel +import asyncio +import gc import json import re import tempfile @@ -277,6 +279,13 @@ class ServingChatTestCase(unittest.TestCase): # to publish one rather than hang the values off a mock manager. reset_context() self.addCleanup(reset_context) + # Tests drive coroutines through get_or_create_event_loop(), which + # creates a fresh loop per call and leaves the previous one unclosed. + # Finalize those loops here, between tests: if the cyclic GC collects + # one mid-import, its ResourceWarning imports tracemalloc while the + # outer import still holds the module-lock bookkeeping, which raises + # KeyError from importlib._bootstrap on Python < 3.12. + self.addCleanup(self._close_event_loops) publish( ServerArgs( model_path="dummy", @@ -314,6 +323,17 @@ class ServingChatTestCase(unittest.TestCase): self.fastapi_request = Mock(spec=Request) self.fastapi_request.headers = {} + @staticmethod + def _close_event_loops(): + try: + loop = asyncio.get_event_loop_policy().get_event_loop() + except RuntimeError: + loop = None + if loop is not None and not loop.is_closed(): + loop.close() + asyncio.set_event_loop(None) + gc.collect() + @staticmethod def _render_tool_results_in_call_order(messages, **kwargs): """Block-level tool_call_id association, like the GLM chat templates.""" diff --git a/test/registered/unit/mem_cache/test_hicache_auto_size.py b/test/registered/unit/mem_cache/test_hicache_auto_size.py new file mode 100644 index 000000000..f0ed470d5 --- /dev/null +++ b/test/registered/unit/mem_cache/test_hicache_auto_size.py @@ -0,0 +1,104 @@ +import unittest +from unittest.mock import Mock, patch + +import torch + +from sglang.srt.arg_groups.overrides import resolution_result +from sglang.srt.mem_cache import hicache_auto_size as sizing +from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool +from sglang.srt.mem_cache.cache_init_params import CacheInitParams +from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool +from sglang.srt.mem_cache.pool_host import base +from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost +from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool +from sglang.srt.runtime_context import get_context, get_memory +from sglang.srt.server_args import ServerArgs +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestHiCacheAutoSize(CustomTestCase): + def test_hybrid_target_draft_uses_its_sidecar_slot_capacity(self): + # EAGLE3 can pair a hybrid target with a plain MHA draft. Full drafts + # follow full target slots; SWA drafts follow the smaller SWA capacity. + target = Mock( + spec=SWAKVPool, + size=128, + full_kv_pool=Mock(size=128, get_kv_size_bytes=Mock(return_value=4096)), + swa_kv_pool=Mock(size=32, get_kv_size_bytes=Mock(return_value=1024)), + ) + draft_mha = Mock( + spec=MHATokenToKVPool, + size=16, + get_kv_size_bytes=Mock(return_value=(128, 128)), + ) + params = CacheInitParams( + disable=False, + req_to_token_pool=None, + token_to_kv_pool_allocator=Mock(get_kvcache=Mock(return_value=target)), + page_size=2, + ) + for draft, expected_sidecar_bytes in ( + (draft_mha, 2048), + (Mock(spec=BaseSWAKVPool, swa_kv_pool=draft_mha), 512), + ): + with self.subTest( + draft_type=type(draft).__name__, bytes=expected_sidecar_bytes + ): + plan = Mock(mode="sidecar", device_pools=(draft,)) + self.assertEqual( + sizing._estimate_hicache_bytes(params, plan), + 4096 + 1024 + expected_sidecar_bytes, + ) + + def test_default_ratio_fits_host_budget_and_pools_book_it(self): + """With only --enable-hierarchical-cache the default ratio shrinks to the + per-rank budget, pools book one snapshot, and an explicit ratio opts out.""" + pool = MHATokenToKVPool( + size=128, + page_size=2, + dtype=torch.float16, + head_num=2, + head_dim=4, + layer_num=2, + device="cpu", + enable_memory_saver=False, + ) + params = CacheInitParams( + disable=False, + req_to_token_pool=None, + token_to_kv_pool_allocator=Mock(get_kvcache=Mock(return_value=pool)), + page_size=2, + ) + rank_budget = 10_000 + # Four ranks per host (e.g. TP8 over two 4-GPU nodes) share what psutil reports. + host_free = base.HICACHE_HOST_MEMORY_RESERVE_BYTES + 4 * rank_budget + with ( + get_context().override_server_args(enable_hierarchical_cache=True), + patch.object(base, "ranks_per_host", return_value=4), + patch.object(base, "available_host_memory_bytes", return_value=host_free), + sizing.auto_size_hicache(params, None, enabled=True), + ): + ratio = get_memory().hicache_ratio + self.assertLess(ratio, 2.0) + host = MHATokenToKVPoolHost( + pool, ratio, 0, 2, "layer_first", pin_memory=False, device="cpu" + ) + self.assertLessEqual(host.size * host.size_per_token, 0.8 * rank_budget) + with self.assertRaisesRegex(ValueError, "Not enough host memory"): + MHATokenToKVPoolHost( + pool, ratio, 0, 2, "layer_first", pin_memory=False, device="cpu" + ) + self.assertIsNone(base._host_memory_budget.get()) + + explicit = ServerArgs(model_path="dummy", hicache_ratio=2.0) + explicit.resolve_once() + self.assertIsNone( + resolution_result(explicit, "hicache_host_memory_fraction", 0.8) + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_host_memory.py b/test/registered/unit/mem_cache/test_host_memory.py new file mode 100644 index 000000000..7e8a9034f --- /dev/null +++ b/test/registered/unit/mem_cache/test_host_memory.py @@ -0,0 +1,130 @@ +"""Exercise container and ancestor budgets using synthetic procfs/cgroup files.""" + +import tempfile +import unittest +from pathlib import Path +from unittest.mock import Mock, patch + +from sglang.srt.mem_cache import host_memory +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +class TestHostMemory(unittest.TestCase): + def setUp(self): + self.temp = tempfile.TemporaryDirectory() + self.addCleanup(self.temp.cleanup) + self.root = Path(self.temp.name) + self.proc = self.root / "proc" + (self.proc / "self").mkdir(parents=True) + self.mount = self.root / "cgroup mount" + self.mount.mkdir() + + def configure(self, membership="/task/engine", mount_root="/", v1=False): + controllers = "memory" if v1 else "" + (self.proc / "self/cgroup").write_text(f"0:{controllers}:{membership}\n") + filesystem = "cgroup" if v1 else "cgroup2" + options = "rw,memory" if v1 else "rw" + escaped = str(self.mount).replace(" ", r"\040") + (self.proc / "self/mountinfo").write_text( + f"1 0 0:1 {mount_root} {escaped} rw - {filesystem} cgroup {options}\n" + ) + + def memory(self, path, usage, maximum="max", high="max", v1=False): + directory = self.mount / path + directory.mkdir(parents=True, exist_ok=True) + files = ( + {"memory.limit_in_bytes": maximum, "memory.usage_in_bytes": usage} + if v1 + else {"memory.max": maximum, "memory.high": high, "memory.current": usage} + ) + for name, value in files.items(): + (directory / name).write_text(str(value)) + + def test_v2_parent_and_high_limits(self): + self.configure() + self.memory("task/engine", 100) + for maximum, high, usage, expected in [ + (1000, "max", 300, 700), + (1000, 600, 300, 300), + ("max", 600, 700, 0), + ("max", "max", 300, None), + ]: + with self.subTest(maximum=maximum, high=high): + self.memory("task", usage, maximum, high) + self.assertEqual( + host_memory._cgroup_memory_headroom(self.proc), expected + ) + self.memory("task", 100, 1000) + self.memory("task/engine", 100, 250) + self.assertEqual(host_memory._cgroup_memory_headroom(self.proc), 150) + + def test_mount_subtree_and_cgroup_namespace(self): + self.memory("engine", 100, 900) + self.memory("", 300, 1000) + for membership in ["/host/task/engine", "/engine"]: + with self.subTest(membership=membership): + self.configure(membership, mount_root="/host/task") + self.assertEqual(host_memory._cgroup_memory_headroom(self.proc), 700) + + def test_v1_parent_limit_and_unlimited_sentinel(self): + self.configure(v1=True) + self.memory("task/engine", 100, 2**63 - 4096, v1=True) + self.memory("task", 400, 1000, v1=True) + self.assertEqual(host_memory._cgroup_memory_headroom(self.proc), 600) + + def test_independent_engines_have_separate_allowances(self): + # Both engines see the same host RAM but have different charged usage. + for task, usage, expected in [("a", 300, 700), ("b", 600, 400)]: + with self.subTest(task=task): + self.configure(f"/{task}/engine") + self.memory(f"{task}/engine", 0) + self.memory(task, usage, 1000) + with ( + patch.object( + host_memory.psutil, + "virtual_memory", + return_value=Mock(available=2000), + ), + patch.object( + host_memory, + "_cgroup_memory_headroom", + return_value=host_memory._cgroup_memory_headroom(self.proc), + ), + ): + self.assertEqual( + host_memory.available_host_memory_bytes(), expected + ) + + def test_host_availability_is_also_a_bound(self): + for cgroup, expected in [(None, 100), (200, 100), (50, 50)]: + with ( + self.subTest(cgroup=cgroup), + patch.object( + host_memory.psutil, + "virtual_memory", + return_value=Mock(available=100), + ), + patch.object( + host_memory, "_cgroup_memory_headroom", return_value=cgroup + ), + ): + self.assertEqual(host_memory.available_host_memory_bytes(), expected) + + def test_unmounted_memory_cgroup_fails(self): + self.configure() + (self.proc / "self/mountinfo").write_text("") + with self.assertRaisesRegex(RuntimeError, "Cannot locate"): + host_memory._cgroup_memory_headroom(self.proc) + + def test_missing_usage_for_known_limit_fails(self): + self.configure() + self.memory("task/engine", 100, 1000) + (self.mount / "task/engine/memory.current").unlink() + with self.assertRaises(FileNotFoundError): + host_memory._cgroup_memory_headroom(self.proc) + + +if __name__ == "__main__": + unittest.main() 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 3d9f39c06..909b8719a 100644 --- a/test/registered/unit/mem_cache/test_mem_pool_host.py +++ b/test/registered/unit/mem_cache/test_mem_pool_host.py @@ -282,11 +282,10 @@ class TestHostMemoryBudget(CustomTestCase): def _budget_with_ranks(self, ranks): # Deliberate single-accessor stub: isolates the budget math from the # topology derivation, which the ranks_per_host case below covers. - fake_mem = unittest.mock.Mock(available=self._AVAILABLE) with ( unittest.mock.patch.object(base, "ranks_per_host", return_value=ranks), unittest.mock.patch.object( - base.psutil, "virtual_memory", return_value=fake_mem + base, "available_host_memory_bytes", return_value=self._AVAILABLE ), ): return base.host_memory_budget_bytes()