diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 926f9f2030..9a0a42fb01 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -905,7 +905,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--disk_cache_dir", type=str, default=None, - help="""Directory used to persist disk cache data. Defaults to a temp directory when not set.""", + help="""Base directory for disk cache. Each server uses an instance-specific subdirectory, + removed after its workers exit. Defaults to a temporary directory when not set.""", ) parser.add_argument( "--enable_dp_prompt_cache_fetch", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 31fa426ef0..e3ea5515a1 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -1,5 +1,6 @@ import multiprocessing as mp import os +import tempfile import uuid import subprocess import math @@ -420,11 +421,17 @@ def _launch_subprocesses(args: StartArgs): if args.enable_cpu_cache: from .multi_level_kv_cache.manager import start_multi_level_kv_cache_manager + instance_disk_cache_dir = None + if args.enable_disk_cache: + cache_base_dir = args.disk_cache_dir or tempfile.gettempdir() + instance_disk_cache_dir = os.path.join(cache_base_dir, f"lightllm_disk_cache_{get_unique_server_name()}") + process_manager.register_disk_cache_dir(instance_disk_cache_dir) + process_manager.start_submodule_processes( start_funcs=[ start_multi_level_kv_cache_manager, ], - start_args=[(args,)], + start_args=[(args, instance_disk_cache_dir)], ) process_manager.start_submodule_processes( diff --git a/lightllm/server/multi_level_kv_cache/disk_cache_worker.py b/lightllm/server/multi_level_kv_cache/disk_cache_worker.py index 542ddbd877..aff400753d 100644 --- a/lightllm/server/multi_level_kv_cache/disk_cache_worker.py +++ b/lightllm/server/multi_level_kv_cache/disk_cache_worker.py @@ -1,12 +1,10 @@ import os -import tempfile import time import math from dataclasses import dataclass -from typing import List, Optional +from typing import List import torch -from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger from .cpu_cache_client import CpuKvCacheClient @@ -37,7 +35,7 @@ def __init__( self, disk_cache_storage_size: float, cpu_cache_client: CpuKvCacheClient, - disk_cache_dir: Optional[str] = None, + cache_dir: str, ): self.cpu_cache_client = cpu_cache_client self._pages_all_idle = False @@ -50,9 +48,6 @@ def __init__( # 读写同时进行时,分配8线程用来写,16线程用来读 max_concurrent_write_tasks = 8 - cache_dir = disk_cache_dir - if not cache_dir: - cache_dir = os.path.join(tempfile.gettempdir(), f"lightllm_disk_cache_{get_unique_server_name()}") os.makedirs(cache_dir, exist_ok=True) cache_file = os.path.join(cache_dir, "cache_file") diff --git a/lightllm/server/multi_level_kv_cache/manager.py b/lightllm/server/multi_level_kv_cache/manager.py index ef5b7369c9..57c679a2ca 100644 --- a/lightllm/server/multi_level_kv_cache/manager.py +++ b/lightllm/server/multi_level_kv_cache/manager.py @@ -27,6 +27,7 @@ class MultiLevelKVCacheManager: def __init__( self, args: StartArgs, + instance_disk_cache_dir, ): self.args: StartArgs = args ports = get_shm_port_args() @@ -56,7 +57,7 @@ def __init__( self.disk_cache_worker = DiskCacheWorker( disk_cache_storage_size=self.args.disk_cache_storage_size, cpu_cache_client=self.cpu_cache_client, - disk_cache_dir=self.args.disk_cache_dir, + cache_dir=instance_disk_cache_dir, ) self.disk_cache_thread = threading.Thread(target=self.disk_cache_worker.run, daemon=True) self.disk_cache_thread.start() @@ -250,7 +251,7 @@ def recv_loop(self): return -def start_multi_level_kv_cache_manager(args, pipe_writer): +def start_multi_level_kv_cache_manager(args, instance_disk_cache_dir, pipe_writer): # 注册graceful 退出的处理 graceful_registry(inspect.currentframe().f_code.co_name) setproctitle.setproctitle(f"lightllm::{get_unique_server_name()}::multi_level_kv_cache") @@ -259,6 +260,7 @@ def start_multi_level_kv_cache_manager(args, pipe_writer): try: manager = MultiLevelKVCacheManager( args=args, + instance_disk_cache_dir=instance_disk_cache_dir, ) except Exception as e: logger.exception(f"start multi_level_kv_cache_manager has exception {str(e)}") diff --git a/lightllm/utils/start_utils.py b/lightllm/utils/start_utils.py index 3d764d7789..4fb2f91f38 100644 --- a/lightllm/utils/start_utils.py +++ b/lightllm/utils/start_utils.py @@ -1,4 +1,5 @@ import os +import shutil import ctypes import signal import subprocess @@ -18,6 +19,10 @@ class SubmoduleManager: def __init__(self): self.processes = [] self.process_names = {} + self.disk_cache_dir = None + + def register_disk_cache_dir(self, cache_dir): + self.disk_cache_dir = cache_dir def start_submodule_processes(self, start_funcs=[], start_args=[]): assert len(start_funcs) == len(start_args) @@ -91,6 +96,15 @@ def terminate_all_processes(self): if alive_pids: logger.warning(f"Processes still alive after SIGKILL: {alive_pids}") + # Only the instance-owned directory may be removed, after its workers exit. + if self.disk_cache_dir is not None and not alive_pids: + try: + shutil.rmtree(self.disk_cache_dir) + except FileNotFoundError: + pass + except OSError: + logger.exception("Failed to remove disk cache directory %s", self.disk_cache_dir) + # recover the gpu compute mode is_enable_mps = get_env_start_args().enable_mps if is_enable_mps: diff --git a/unit_tests/server/multi_level_kv_cache/test_disk_cache_lifecycle.py b/unit_tests/server/multi_level_kv_cache/test_disk_cache_lifecycle.py new file mode 100644 index 0000000000..d994063ee8 --- /dev/null +++ b/unit_tests/server/multi_level_kv_cache/test_disk_cache_lifecycle.py @@ -0,0 +1,43 @@ +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + +from lightllm.utils import start_utils + + +@pytest.mark.parametrize("worker_alive", [False, True]) +def test_shutdown_only_removes_instance_directory_after_workers_exit(tmp_path, monkeypatch, worker_alive): + base = tmp_path / "cache" + instance = base / "lightllm_disk_cache_current" + other = base / "lightllm_disk_cache_other" + instance.mkdir(parents=True) + other.mkdir() + manager = start_utils.SubmoduleManager() + manager.register_disk_cache_dir(str(instance)) + worker = Mock(pid=12345) + manager.processes = [worker] + monkeypatch.setattr(start_utils, "kill_recursive", Mock()) + monkeypatch.setattr( + start_utils.psutil, "wait_procs", lambda *_args, **_kwargs: ([], [worker] if worker_alive else []) + ) + monkeypatch.setattr(start_utils, "is_process_active", lambda _pid: worker_alive) + monkeypatch.setattr("lightllm.utils.envs_utils.get_env_start_args", lambda: SimpleNamespace(enable_mps=False)) + + manager.terminate_all_processes() + assert instance.exists() == worker_alive + assert other.exists() and base.exists() + if not worker_alive: + manager.terminate_all_processes() + + +def test_disk_worker_uses_exact_instance_directory(tmp_path, monkeypatch): + import torch + from lightllm.server.multi_level_kv_cache import disk_cache_worker + + service = Mock(return_value=SimpleNamespace(_n=1)) + monkeypatch.setattr(disk_cache_worker, "PyLocalCacheService", service) + directory = tmp_path / "lightllm_disk_cache_current" + disk_cache_worker.DiskCacheWorker(1, SimpleNamespace(cpu_kv_cache_tensor=torch.zeros((2, 4))), str(directory)) + assert directory.is_dir() + assert service.call_args.kwargs["file"] == str(directory / "cache_file")