From 34a8641653cbe4f9d429861fc841fa757171d1df Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 25 Aug 2026 15:20:35 +0800 Subject: [PATCH 1/3] fix(pd): reserve transfer pages in memory profile --- .../deepseek2_mem_manager.py | 11 +- .../kv_cache_mem_manager/mem_manager.py | 35 +++- .../common/kv_cache_mem_manager/__init__.py | 1 + .../test_pd_memory_budget.py | 197 ++++++++++++++++++ 4 files changed, 231 insertions(+), 13 deletions(-) create mode 100644 unit_tests/common/kv_cache_mem_manager/__init__.py create mode 100644 unit_tests/common/kv_cache_mem_manager/test_pd_memory_budget.py diff --git a/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py index 9eb02b963c..adf81a14d8 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py @@ -12,7 +12,6 @@ class Deepseek2MemoryManager(MemoryManager): - operator_class = Deepseek2MemOperator def __init__(self, size, dtype, head_num, head_dim, layer_num, always_copy=False, mem_fraction=0.9): @@ -28,14 +27,8 @@ def get_cell_size(self): def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): self.kv_buffer = torch.empty((layer_num, size + 1, head_num, head_dim), dtype=dtype, device="cuda") - def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: - self.kv_move_buffer = torch.empty( - (page_num, page_size, self.layer_num, self.head_num, self.head_dim), dtype=self.dtype, device="cuda" - ) - self._buffer_mem_indexes_tensors = [ - torch.empty((page_size,), dtype=torch.int64, device="cpu", pin_memory=True) for _ in range(page_num) - ] - return self.kv_move_buffer + def get_paged_kv_move_buffer_shape(self, page_num, page_size): + return (page_num, page_size, self.layer_num, self.head_num, self.head_dim) def write_mem_to_page_kv_move_buffer( self, diff --git a/lightllm/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py index 658d3e899c..9f3bea73e4 100755 --- a/lightllm/common/kv_cache_mem_manager/mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py @@ -1,5 +1,6 @@ import re import os +import math import torch import torch.distributed as dist import torch.multiprocessing as mp @@ -23,7 +24,6 @@ class MemoryManager: - operator_class = NormalMemOperator def __init__(self, size, dtype, head_num, head_dim, layer_num, always_copy=False, mem_fraction=0.9): @@ -58,6 +58,14 @@ def get_att_input_params(self, layer_index: int) -> Tuple[Any, Any]: def get_cell_size(self): return 2 * self.head_num * self.head_dim * self.layer_num * torch._utils._element_size(self.dtype) + def get_paged_kv_move_buffer_shape(self, page_num, page_size): + num_kv_head = get_num_key_value_heads(get_env_start_args().model_dir) + return (page_num, page_size, self.layer_num, 2 * num_kv_head, self.head_dim) + + def get_paged_kv_move_buffer_size_in_bytes(self, page_num, page_size): + shape = self.get_paged_kv_move_buffer_shape(page_num, page_size) + return math.prod(shape) * torch._utils._element_size(self.dtype) + def profile_size(self, mem_fraction): if self.size is not None: return @@ -65,14 +73,34 @@ def profile_size(self, mem_fraction): torch.cuda.empty_cache() world_size = dist.get_world_size() available_memory = get_available_gpu_memory(world_size) - get_total_gpu_memory() * (1 - mem_fraction) + args = get_env_start_args() + pd_kv_move_buffer_size_in_bytes = 0 + if args.run_mode in ["prefill", "decode"]: + pd_kv_move_buffer_size_in_bytes = self.get_paged_kv_move_buffer_size_in_bytes( + page_num=args.pd_kv_page_num, + page_size=args.pd_kv_page_size, + ) + available_memory -= pd_kv_move_buffer_size_in_bytes / 1024 ** 3 cell_size = self.get_cell_size() self.size = int(available_memory * 1024 ** 3 / cell_size) if world_size > 1: tensor = torch.tensor(self.size, dtype=torch.int64, device=f"cuda:{get_current_device_id()}") dist.all_reduce(tensor, op=dist.ReduceOp.MIN) self.size = tensor.item() + if pd_kv_move_buffer_size_in_bytes > 0 and self.size <= 0: + raise RuntimeError( + "PD KV transfer page buffer reservation leaves no memory for the token KV cache; " + "reduce --pd_kv_page_size or --pd_kv_page_num, or increase --mem_fraction" + ) + pd_kv_move_buffer_log = "" + if pd_kv_move_buffer_size_in_bytes > 0: + pd_kv_move_buffer_log = ( + f"{str(pd_kv_move_buffer_size_in_bytes / 1024 ** 3)} GB is reserved " + "for the PD KV transfer page buffer\n" + ) logger.info( - f"{str(available_memory)} GB space is available after load the model weight\n" + f"{str(available_memory)} GB space is available for the token KV cache\n" + f"{pd_kv_move_buffer_log}" f"{str(cell_size / 1024 ** 2)} MB is the size of one token kv cache\n" f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n" ) @@ -86,9 +114,8 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): self.kv_buffer = torch.empty((layer_num, size + 1, 2 * head_num, head_dim), dtype=dtype, device="cuda") def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: - num_kv_head = get_num_key_value_heads(get_env_start_args().model_dir) self.kv_move_buffer = torch.empty( - (page_num, page_size, self.layer_num, 2 * num_kv_head, self.head_dim), dtype=self.dtype, device="cuda" + self.get_paged_kv_move_buffer_shape(page_num, page_size), dtype=self.dtype, device="cuda" ) self._buffer_mem_indexes_tensors = [ torch.empty((page_size,), dtype=torch.int64, device="cpu", pin_memory=True) for _ in range(page_num) diff --git a/unit_tests/common/kv_cache_mem_manager/__init__.py b/unit_tests/common/kv_cache_mem_manager/__init__.py new file mode 100644 index 0000000000..8b13789179 --- /dev/null +++ b/unit_tests/common/kv_cache_mem_manager/__init__.py @@ -0,0 +1 @@ + diff --git a/unit_tests/common/kv_cache_mem_manager/test_pd_memory_budget.py b/unit_tests/common/kv_cache_mem_manager/test_pd_memory_budget.py new file mode 100644 index 0000000000..a649dd7b2c --- /dev/null +++ b/unit_tests/common/kv_cache_mem_manager/test_pd_memory_budget.py @@ -0,0 +1,197 @@ +import importlib.util +import sys +import types +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch + + +@pytest.fixture(scope="module") +def manager_modules(): + """Import GPU modules without leaking local macOS compatibility stubs.""" + lightllm_modules_before = {name for name in sys.modules if name.startswith("lightllm")} + module_patch = nullcontext() + if importlib.util.find_spec("triton") is None: + triton = types.ModuleType("triton") + triton.jit = lambda fn: fn + triton.cdiv = lambda value, divisor: (value + divisor - 1) // divisor + triton_language = types.ModuleType("triton.language") + triton_language.constexpr = object() + triton.language = triton_language + transformers = types.ModuleType("transformers") + transformers.AutoModelForCausalLM = object() + module_patch = patch.dict( + sys.modules, + { + "transformers": transformers, + "triton": triton, + "triton.language": triton_language, + }, + ) + + try: + with module_patch: + from lightllm.common.kv_cache_mem_manager import Deepseek2MemoryManager, MemoryManager + from lightllm.common.kv_cache_mem_manager import mem_manager as mem_manager_module + + yield SimpleNamespace( + deepseek_class=Deepseek2MemoryManager, + memory_manager_class=MemoryManager, + module=mem_manager_module, + ) + finally: + for module_name in list(sys.modules): + if module_name.startswith("lightllm") and module_name not in lightllm_modules_before: + sys.modules.pop(module_name, None) + + +def _profile_manager(monkeypatch, manager_modules, manager, *, run_mode, mem_fraction=0.8, page_size=4096): + mem_manager_module = manager_modules.module + monkeypatch.setattr(mem_manager_module.torch.cuda, "empty_cache", lambda: None) + monkeypatch.setattr(mem_manager_module.dist, "get_world_size", lambda: 1) + monkeypatch.setattr(mem_manager_module, "get_available_gpu_memory", lambda world_size: 10.0) + monkeypatch.setattr(mem_manager_module, "get_total_gpu_memory", lambda: 10.0) + monkeypatch.setattr( + mem_manager_module, + "get_env_start_args", + lambda: SimpleNamespace( + run_mode=run_mode, + model_dir="unused", + pd_kv_page_num=16, + pd_kv_page_size=page_size, + ), + ) + + manager.profile_size(mem_fraction) + + +@pytest.mark.parametrize("run_mode", ["prefill", "decode"]) +@pytest.mark.parametrize("page_size, expected_token_capacity", [(1024, 427911), (4096, 231303)]) +def test_pd_profile_reserves_qwen35_transfer_page_buffer( + monkeypatch, manager_modules, run_mode, page_size, expected_token_capacity +): + manager = object.__new__(manager_modules.memory_manager_class) + manager.size = None + manager.dtype = torch.bfloat16 + manager.head_num = 1 + manager.head_dim = 256 + manager.layer_num = 17 + monkeypatch.setattr(manager_modules.module, "get_num_key_value_heads", lambda model_dir: 4) + + _profile_manager(monkeypatch, manager_modules, manager, run_mode=run_mode, page_size=page_size) + + # At page_size=4096, mem_fraction leaves 8 GiB and the Qwen3.5 PD page buffer + # occupies 4.25 GiB: 16 * 4096 * 17 layers * 2 K/V * 4 global heads * 256 * 2 bytes. + assert manager.size == expected_token_capacity + + +@pytest.mark.parametrize("run_mode", ["normal", "pd_master"]) +def test_non_pd_worker_profile_does_not_reserve_pd_transfer_buffer(monkeypatch, manager_modules, run_mode): + manager = object.__new__(manager_modules.memory_manager_class) + manager.size = None + manager.dtype = torch.bfloat16 + manager.head_num = 1 + manager.head_dim = 256 + manager.layer_num = 17 + monkeypatch.setattr(manager_modules.module, "get_num_key_value_heads", lambda model_dir: 4) + + _profile_manager(monkeypatch, manager_modules, manager, run_mode=run_mode) + + assert manager.size == 493447 + + +def test_pd_profile_uses_deepseek_transfer_buffer_layout(monkeypatch, manager_modules): + manager = object.__new__(manager_modules.deepseek_class) + manager.size = None + manager.dtype = torch.bfloat16 + manager.head_num = 1 + manager.head_dim = 576 + manager.layer_num = 61 + + _profile_manager(monkeypatch, manager_modules, manager, run_mode="decode") + + # MLA pages store one compressed KV tensor per local head rather than K/V + # tensors for every global KV head. + assert manager.size == 56702 + + +def test_explicit_token_capacity_skips_pd_memory_profile(monkeypatch, manager_modules): + manager = object.__new__(manager_modules.memory_manager_class) + manager.size = 12345 + + monkeypatch.setattr( + manager_modules.module, + "get_env_start_args", + lambda: pytest.fail("explicit capacity must not inspect PD profile arguments"), + ) + + manager.profile_size(mem_fraction=0.8) + + assert manager.size == 12345 + + +@pytest.mark.parametrize("mem_fraction", [0.4, 0.425001]) +def test_pd_profile_fails_when_transfer_buffer_leaves_no_token_capacity(monkeypatch, manager_modules, mem_fraction): + manager = object.__new__(manager_modules.memory_manager_class) + manager.size = None + manager.dtype = torch.bfloat16 + manager.head_num = 1 + manager.head_dim = 256 + manager.layer_num = 17 + monkeypatch.setattr(manager_modules.module, "get_num_key_value_heads", lambda model_dir: 4) + + with pytest.raises( + RuntimeError, + match=r"reduce --pd_kv_page_size or --pd_kv_page_num, or increase --mem_fraction", + ): + _profile_manager(monkeypatch, manager_modules, manager, run_mode="decode", mem_fraction=mem_fraction) + + +@pytest.mark.parametrize( + "manager_name, head_num, head_dim, layer_num, global_kv_heads, expected_shape", + [ + ("base", 1, 256, 17, 4, (2, 32, 17, 8, 256)), + ("deepseek", 1, 576, 61, 4, (2, 32, 61, 1, 576)), + ], +) +def test_pd_transfer_allocation_uses_profiled_buffer_shape( + monkeypatch, + manager_modules, + manager_name, + head_num, + head_dim, + layer_num, + global_kv_heads, + expected_shape, +): + manager_class = manager_modules.memory_manager_class if manager_name == "base" else manager_modules.deepseek_class + manager = object.__new__(manager_class) + manager.dtype = torch.bfloat16 + manager.head_num = head_num + manager.head_dim = head_dim + manager.layer_num = layer_num + monkeypatch.setattr( + manager_modules.module, + "get_env_start_args", + lambda: SimpleNamespace(model_dir="unused"), + ) + monkeypatch.setattr( + manager_modules.module, + "get_num_key_value_heads", + lambda model_dir: global_kv_heads, + ) + real_empty = torch.empty + + def cpu_empty(*args, **kwargs): + kwargs["device"] = "cpu" + kwargs.pop("pin_memory", None) + return real_empty(*args, **kwargs) + + monkeypatch.setattr(torch, "empty", cpu_empty) + + move_buffer = manager.alloc_paged_kv_move_buffer(page_num=2, page_size=32) + + assert tuple(move_buffer.shape) == expected_shape From a9779224791aa94848e97e655aae72c74a359e45 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 25 Aug 2026 15:23:18 +0800 Subject: [PATCH 2/3] style(tests): keep package marker empty --- unit_tests/common/kv_cache_mem_manager/__init__.py | 1 - 1 file changed, 1 deletion(-) diff --git a/unit_tests/common/kv_cache_mem_manager/__init__.py b/unit_tests/common/kv_cache_mem_manager/__init__.py index 8b13789179..e69de29bb2 100644 --- a/unit_tests/common/kv_cache_mem_manager/__init__.py +++ b/unit_tests/common/kv_cache_mem_manager/__init__.py @@ -1 +0,0 @@ - From 7bc5763c8d5b6712dcfc4da4cc1f87efa368f1a5 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 25 Aug 2026 16:46:05 +0800 Subject: [PATCH 3/3] chore(tests): remove PD memory budget tests --- .../common/kv_cache_mem_manager/__init__.py | 0 .../test_pd_memory_budget.py | 197 ------------------ 2 files changed, 197 deletions(-) delete mode 100644 unit_tests/common/kv_cache_mem_manager/__init__.py delete mode 100644 unit_tests/common/kv_cache_mem_manager/test_pd_memory_budget.py diff --git a/unit_tests/common/kv_cache_mem_manager/__init__.py b/unit_tests/common/kv_cache_mem_manager/__init__.py deleted file mode 100644 index e69de29bb2..0000000000 diff --git a/unit_tests/common/kv_cache_mem_manager/test_pd_memory_budget.py b/unit_tests/common/kv_cache_mem_manager/test_pd_memory_budget.py deleted file mode 100644 index a649dd7b2c..0000000000 --- a/unit_tests/common/kv_cache_mem_manager/test_pd_memory_budget.py +++ /dev/null @@ -1,197 +0,0 @@ -import importlib.util -import sys -import types -from contextlib import nullcontext -from types import SimpleNamespace -from unittest.mock import patch - -import pytest -import torch - - -@pytest.fixture(scope="module") -def manager_modules(): - """Import GPU modules without leaking local macOS compatibility stubs.""" - lightllm_modules_before = {name for name in sys.modules if name.startswith("lightllm")} - module_patch = nullcontext() - if importlib.util.find_spec("triton") is None: - triton = types.ModuleType("triton") - triton.jit = lambda fn: fn - triton.cdiv = lambda value, divisor: (value + divisor - 1) // divisor - triton_language = types.ModuleType("triton.language") - triton_language.constexpr = object() - triton.language = triton_language - transformers = types.ModuleType("transformers") - transformers.AutoModelForCausalLM = object() - module_patch = patch.dict( - sys.modules, - { - "transformers": transformers, - "triton": triton, - "triton.language": triton_language, - }, - ) - - try: - with module_patch: - from lightllm.common.kv_cache_mem_manager import Deepseek2MemoryManager, MemoryManager - from lightllm.common.kv_cache_mem_manager import mem_manager as mem_manager_module - - yield SimpleNamespace( - deepseek_class=Deepseek2MemoryManager, - memory_manager_class=MemoryManager, - module=mem_manager_module, - ) - finally: - for module_name in list(sys.modules): - if module_name.startswith("lightllm") and module_name not in lightllm_modules_before: - sys.modules.pop(module_name, None) - - -def _profile_manager(monkeypatch, manager_modules, manager, *, run_mode, mem_fraction=0.8, page_size=4096): - mem_manager_module = manager_modules.module - monkeypatch.setattr(mem_manager_module.torch.cuda, "empty_cache", lambda: None) - monkeypatch.setattr(mem_manager_module.dist, "get_world_size", lambda: 1) - monkeypatch.setattr(mem_manager_module, "get_available_gpu_memory", lambda world_size: 10.0) - monkeypatch.setattr(mem_manager_module, "get_total_gpu_memory", lambda: 10.0) - monkeypatch.setattr( - mem_manager_module, - "get_env_start_args", - lambda: SimpleNamespace( - run_mode=run_mode, - model_dir="unused", - pd_kv_page_num=16, - pd_kv_page_size=page_size, - ), - ) - - manager.profile_size(mem_fraction) - - -@pytest.mark.parametrize("run_mode", ["prefill", "decode"]) -@pytest.mark.parametrize("page_size, expected_token_capacity", [(1024, 427911), (4096, 231303)]) -def test_pd_profile_reserves_qwen35_transfer_page_buffer( - monkeypatch, manager_modules, run_mode, page_size, expected_token_capacity -): - manager = object.__new__(manager_modules.memory_manager_class) - manager.size = None - manager.dtype = torch.bfloat16 - manager.head_num = 1 - manager.head_dim = 256 - manager.layer_num = 17 - monkeypatch.setattr(manager_modules.module, "get_num_key_value_heads", lambda model_dir: 4) - - _profile_manager(monkeypatch, manager_modules, manager, run_mode=run_mode, page_size=page_size) - - # At page_size=4096, mem_fraction leaves 8 GiB and the Qwen3.5 PD page buffer - # occupies 4.25 GiB: 16 * 4096 * 17 layers * 2 K/V * 4 global heads * 256 * 2 bytes. - assert manager.size == expected_token_capacity - - -@pytest.mark.parametrize("run_mode", ["normal", "pd_master"]) -def test_non_pd_worker_profile_does_not_reserve_pd_transfer_buffer(monkeypatch, manager_modules, run_mode): - manager = object.__new__(manager_modules.memory_manager_class) - manager.size = None - manager.dtype = torch.bfloat16 - manager.head_num = 1 - manager.head_dim = 256 - manager.layer_num = 17 - monkeypatch.setattr(manager_modules.module, "get_num_key_value_heads", lambda model_dir: 4) - - _profile_manager(monkeypatch, manager_modules, manager, run_mode=run_mode) - - assert manager.size == 493447 - - -def test_pd_profile_uses_deepseek_transfer_buffer_layout(monkeypatch, manager_modules): - manager = object.__new__(manager_modules.deepseek_class) - manager.size = None - manager.dtype = torch.bfloat16 - manager.head_num = 1 - manager.head_dim = 576 - manager.layer_num = 61 - - _profile_manager(monkeypatch, manager_modules, manager, run_mode="decode") - - # MLA pages store one compressed KV tensor per local head rather than K/V - # tensors for every global KV head. - assert manager.size == 56702 - - -def test_explicit_token_capacity_skips_pd_memory_profile(monkeypatch, manager_modules): - manager = object.__new__(manager_modules.memory_manager_class) - manager.size = 12345 - - monkeypatch.setattr( - manager_modules.module, - "get_env_start_args", - lambda: pytest.fail("explicit capacity must not inspect PD profile arguments"), - ) - - manager.profile_size(mem_fraction=0.8) - - assert manager.size == 12345 - - -@pytest.mark.parametrize("mem_fraction", [0.4, 0.425001]) -def test_pd_profile_fails_when_transfer_buffer_leaves_no_token_capacity(monkeypatch, manager_modules, mem_fraction): - manager = object.__new__(manager_modules.memory_manager_class) - manager.size = None - manager.dtype = torch.bfloat16 - manager.head_num = 1 - manager.head_dim = 256 - manager.layer_num = 17 - monkeypatch.setattr(manager_modules.module, "get_num_key_value_heads", lambda model_dir: 4) - - with pytest.raises( - RuntimeError, - match=r"reduce --pd_kv_page_size or --pd_kv_page_num, or increase --mem_fraction", - ): - _profile_manager(monkeypatch, manager_modules, manager, run_mode="decode", mem_fraction=mem_fraction) - - -@pytest.mark.parametrize( - "manager_name, head_num, head_dim, layer_num, global_kv_heads, expected_shape", - [ - ("base", 1, 256, 17, 4, (2, 32, 17, 8, 256)), - ("deepseek", 1, 576, 61, 4, (2, 32, 61, 1, 576)), - ], -) -def test_pd_transfer_allocation_uses_profiled_buffer_shape( - monkeypatch, - manager_modules, - manager_name, - head_num, - head_dim, - layer_num, - global_kv_heads, - expected_shape, -): - manager_class = manager_modules.memory_manager_class if manager_name == "base" else manager_modules.deepseek_class - manager = object.__new__(manager_class) - manager.dtype = torch.bfloat16 - manager.head_num = head_num - manager.head_dim = head_dim - manager.layer_num = layer_num - monkeypatch.setattr( - manager_modules.module, - "get_env_start_args", - lambda: SimpleNamespace(model_dir="unused"), - ) - monkeypatch.setattr( - manager_modules.module, - "get_num_key_value_heads", - lambda model_dir: global_kv_heads, - ) - real_empty = torch.empty - - def cpu_empty(*args, **kwargs): - kwargs["device"] = "cpu" - kwargs.pop("pin_memory", None) - return real_empty(*args, **kwargs) - - monkeypatch.setattr(torch, "empty", cpu_empty) - - move_buffer = manager.alloc_paged_kv_move_buffer(page_num=2, page_size=32) - - assert tuple(move_buffer.shape) == expected_shape