Skip to content
Closed
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
3 changes: 2 additions & 1 deletion lightllm/server/api_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
9 changes: 8 additions & 1 deletion lightllm/server/api_start.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import multiprocessing as mp
import os
import tempfile
import uuid
import subprocess
import math
Expand Down Expand Up @@ -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(
Expand Down
9 changes: 2 additions & 7 deletions lightllm/server/multi_level_kv_cache/disk_cache_worker.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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
Expand All @@ -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")

Expand Down
6 changes: 4 additions & 2 deletions lightllm/server/multi_level_kv_cache/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ class MultiLevelKVCacheManager:
def __init__(
self,
args: StartArgs,
instance_disk_cache_dir,
):
self.args: StartArgs = args
ports = get_shm_port_args()
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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")
Expand All @@ -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)}")
Expand Down
14 changes: 14 additions & 0 deletions lightllm/utils/start_utils.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import os
import shutil
import ctypes
import signal
import subprocess
Expand All @@ -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)
Expand Down Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
@@ -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")
Loading