From 65bde0abc0d28a59e22622000d77b9fe1a19ed4a Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 31 Aug 2026 13:50:41 +0800 Subject: [PATCH 1/9] feat(pd): prioritize cache-friendly requests during admission --- lightllm/server/api_cli.py | 19 +++ lightllm/server/core/objs/start_args_type.py | 2 + .../httpserver_for_pd_master/manager.py | 28 ++++- .../pd_selector/cache_aware.py | 12 ++ .../pd_selector/pd_selector.py | 11 +- .../test_pd_master_cached_tokens.py | 1 + unit_tests/server/test_pd_cache_aware.py | 11 ++ unit_tests/server/test_pd_master_mode.py | 113 ++++++++++++++++++ 8 files changed, 195 insertions(+), 2 deletions(-) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 60b5fad4e..447fa2e5e 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -73,6 +73,25 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: action="store_true", help="Disable PD master admission control based on the total capacity of registered decode nodes.", ) + parser.add_argument( + "--pd_master_decode_waiting_queue_ratio", + type=float, + default=0.5, + help=( + "Extra PD master waiting queue size as a ratio of the total capacity of registered decode nodes. " + "For example, 0.5 allows an additional waiting queue equal to 50%% of decode capacity. Default: 0.5." + ), + ) + parser.add_argument( + "--pd_master_cache_aware_queue_reserved_ratio", + type=float, + default=0.5, + help=( + "Fraction of the PD master waiting queue reserved for requests with reusable prompt cache. " + "Cache-aware admission gradually unlocks this reserved capacity based on the estimated cache hit rate. " + "Set to 0 to disable cache-aware reservation. Default: 0.5." + ), + ) parser.add_argument( "--pd_trans_mode", type=str, diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index a9aef608b..0f14e53d8 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -23,6 +23,8 @@ class StartArgs: pd_master_port: int = field(default=1212) pd_master_mode: str = field(default="elastic") disable_pd_master_decode_capacity_limit: bool = field(default=False) + pd_master_decode_waiting_queue_ratio: float = field(default=0.5) + pd_master_cache_aware_queue_reserved_ratio: float = field(default=0.5) pd_trans_mode: str = field(default="nccl", metadata={"choices": ["nccl", "nixl"]}) config_server_host: str = field(default=None) config_server_port: int = field(default=None) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index b422bf770..4ac9ffc3f 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -4,6 +4,7 @@ import uvloop import time import datetime +import math import ujson as json import pickle import httpx @@ -36,6 +37,13 @@ def __init__( args: StartArgs, ): self.args = args + if ( + not math.isfinite(args.pd_master_decode_waiting_queue_ratio) + or args.pd_master_decode_waiting_queue_ratio < 0 + ): + raise ValueError("pd_master_decode_waiting_queue_ratio must be a non-negative finite number") + if not 0 <= args.pd_master_cache_aware_queue_reserved_ratio <= 1: + raise ValueError("pd_master_cache_aware_queue_reserved_ratio must be between 0 and 1") self.max_req_total_len = args.max_req_total_len assert self.max_req_total_len is not None self.metric_client = MetricClient(get_shm_port_args().metric_port) @@ -131,9 +139,27 @@ async def generate( ): if not self.args.disable_pd_master_decode_capacity_limit: decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.pd_manager.decode_nodes) - if self.running_request_count >= decode_capacity: + waiting_queue_capacity = math.ceil(decode_capacity * self.args.pd_master_decode_waiting_queue_ratio) + hard_admission_limit = decode_capacity + waiting_queue_capacity + if self.running_request_count >= hard_admission_limit: raise ServerBusyError() + reserved_ratio = self.args.pd_master_cache_aware_queue_reserved_ratio + general_waiting_capacity = math.ceil(waiting_queue_capacity * (1.0 - reserved_ratio)) + general_admission_limit = decode_capacity + general_waiting_capacity + if self.running_request_count >= general_admission_limit: + estimate_cache_hit_rate = getattr(self.pd_manager.selector, "estimate_prompt_cache_hit_rate", None) + estimated_cache_hit_rate = estimate_cache_hit_rate(prompt) if estimate_cache_hit_rate else None + if estimated_cache_hit_rate is not None: + if not math.isfinite(estimated_cache_hit_rate): + estimated_cache_hit_rate = 0.0 + estimated_cache_hit_rate = min(max(estimated_cache_hit_rate, 0.0), 1.0) + request_waiting_capacity = math.ceil( + waiting_queue_capacity * (1.0 - reserved_ratio * (1.0 - estimated_cache_hit_rate)) + ) + if self.running_request_count >= decode_capacity + request_waiting_capacity: + raise ServerBusyError() + was_idle = self.running_request_count == 0 self.running_request_count += 1 if was_idle: diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index adf911bc2..08d7a6e90 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -156,6 +156,18 @@ def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: self.balance_rel_threshold_controller.append(cache_hit_rate) self.balance_rel_threshold_controller.update_config(self.config) + def estimate_cache_hit_rate(self, workers: List[PD_Client_Obj], request_text: str) -> float: + """Estimate reusable prompt cache on currently connected prefill workers.""" + if not workers or not request_text: + return 0.0 + + result = self.prompt_cache_tree.prefix_match(request_text) + if result.prefill_node is None or not any(worker.client_ip_port == result.prefill_node for worker in workers): + return 0.0 + if result.input_char_count == 0: + return 0.0 + return min(max(result.matched_char_count / result.input_char_count, 0.0), 1.0) + def _get_cache_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Optional[PD_Client_Obj]: """在指定候选节点中返回达到匹配阈值的 cache 节点。""" result = self.prompt_cache_tree.prefix_match(request_text) diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py index 5474806b7..a8c3a6e04 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py @@ -1,5 +1,5 @@ import random -from typing import Union, List, Tuple, Dict +from typing import Union, List, Tuple, Dict, Optional from lightllm.server.pd_io_struct import PD_Client_Obj from lightllm.server.core.objs import SamplingParams from lightllm.server.multimodal_params import MultimodalParams @@ -30,6 +30,10 @@ def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: """记录推理侧返回的 prompt cache 命中率;非 cache-aware 策略无需处理。""" return + def estimate_prompt_cache_hit_rate(self, prompt: Union[str, List[int]]) -> Optional[float]: + """Return None when this selector cannot estimate reusable prompt cache.""" + return None + class RandomSelector(PDSelector): """随机选择器""" @@ -100,3 +104,8 @@ def select_p_d_node( def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: self.policy.record_prompt_cache_hit_rate(cache_hit_rate) + + def estimate_prompt_cache_hit_rate(self, prompt: Union[str, List[int]]) -> Optional[float]: + if not isinstance(prompt, str): + return 0.0 + return self.policy.estimate_cache_hit_rate(self.prefill_nodes, prompt) diff --git a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py index a3dc268a9..da07759a9 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -15,6 +15,7 @@ def _make_manager(monkeypatch): ) monkeypatch.setattr(SamplingParams, "from_buffer_copy", classmethod(lambda cls, other: copy.copy(other))) mgr = object.__new__(HttpServerManagerForPDMaster) + mgr.args = SimpleNamespace(disable_pd_master_decode_capacity_limit=True) mgr.running_request_count = 0 counter = [0] diff --git a/unit_tests/server/test_pd_cache_aware.py b/unit_tests/server/test_pd_cache_aware.py index 6bdde7857..45fe09173 100644 --- a/unit_tests/server/test_pd_cache_aware.py +++ b/unit_tests/server/test_pd_cache_aware.py @@ -73,6 +73,17 @@ def test_cache_aware_updates_threshold_from_inference_cache_hit_rate(): assert policy.config.balance_rel_threshold == pytest.approx(1.55) +def test_cache_aware_estimates_hit_rate_only_for_connected_worker(): + policy = CacheAwarePolicy(CacheAwareConfig(sample_stride=1)) + cached_worker = _worker("10.0.0.1:8000") + prompt = "shared conversation history and a new user turn" + policy.prompt_cache_tree.insert(prompt[:-10], cached_worker.client_ip_port) + + expected_hit_rate = len(prompt[:-10]) / len(prompt) + assert policy.estimate_cache_hit_rate([cached_worker], prompt) == pytest.approx(expected_hit_rate) + assert policy.estimate_cache_hit_rate([_worker("10.0.0.2:8000")], prompt) == 0.0 + + def test_cache_aware_keeps_cache_worker_when_inflight_load_is_balanced(): policy = CacheAwarePolicy() cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=110, dispatched_req_num=2) diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index edb758a94..fed4563f0 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -5,8 +5,10 @@ import pytest from easydict import EasyDict +from lightllm.server.api_cli import make_argument_parser from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager +from lightllm.utils.error_utils import ServerBusyError def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): @@ -85,6 +87,22 @@ def test_pd_master_models_endpoint_has_created_timestamp(monkeypatch): assert response.data[0].created == 1234 +def test_pd_master_decode_waiting_queue_ratio_cli_default_and_override(): + default_args = make_argument_parser().parse_args([]) + assert default_args.pd_master_decode_waiting_queue_ratio == 0.5 + assert default_args.pd_master_cache_aware_queue_reserved_ratio == 0.5 + configured_args = make_argument_parser().parse_args( + [ + "--pd_master_decode_waiting_queue_ratio", + "0.25", + "--pd_master_cache_aware_queue_reserved_ratio", + "0.75", + ] + ) + assert configured_args.pd_master_decode_waiting_queue_ratio == 0.25 + assert configured_args.pd_master_cache_aware_queue_reserved_ratio == 0.75 + + def test_elastic_pd_nodes_are_ready_with_at_least_one_node_of_each_role(): manager = PDManager(StartArgs(pd_master_mode="elastic")) assert manager.is_pd_nodes_ready() is False @@ -262,12 +280,106 @@ def test_pd_master_inference_health_matches_normal_node_semantics(monkeypatch): assert HttpServerManagerForPDMaster.is_healthy(manager) is True +@pytest.mark.parametrize( + ("waiting_queue_ratio", "running_request_count", "is_rejected"), + [ + (0.5, 11, False), + (0.5, 12, True), + (0.25, 9, False), + (0.25, 10, True), + (0.0, 7, False), + (0.0, 8, True), + ], +) +def test_pd_master_admission_limit_includes_configurable_waiting_queue( + waiting_queue_ratio, running_request_count, is_rejected +): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs(pd_master_decode_waiting_queue_ratio=waiting_queue_ratio) + manager.pd_manager = SimpleNamespace( + decode_nodes=[ + SimpleNamespace(start_args={"running_max_req_size": 3}), + SimpleNamespace(start_args={"running_max_req_size": 5}), + ], + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: None), + ) + manager.running_request_count = running_request_count + manager.latest_success_infer_time = 0 + + async def fake_generate(*_args): + yield "result" + + manager._generate = fake_generate + + async def consume_one_result(): + generator = manager.generate(None, None, None, None) + try: + assert await generator.__anext__() == "result" + finally: + await generator.aclose() + + if is_rejected: + with pytest.raises(ServerBusyError): + asyncio.run(consume_one_result()) + assert manager.running_request_count == running_request_count + else: + asyncio.run(consume_one_result()) + assert manager.running_request_count == running_request_count + + +@pytest.mark.parametrize( + ("estimated_cache_hit_rate", "running_request_count", "is_rejected"), + [ + (0.0, 9, False), + (0.0, 10, True), + (0.5, 10, False), + (0.5, 11, True), + (1.0, 11, False), + (1.0, 12, True), + (None, 11, False), + (None, 12, True), + ], +) +def test_pd_master_admission_reserves_queue_for_cache_friendly_requests( + estimated_cache_hit_rate, running_request_count, is_rejected +): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs() + manager.pd_manager = SimpleNamespace( + decode_nodes=[SimpleNamespace(start_args={"running_max_req_size": 8})], + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: estimated_cache_hit_rate), + ) + manager.running_request_count = running_request_count + manager.latest_success_infer_time = 0 + + async def fake_generate(*_args): + yield "result" + + manager._generate = fake_generate + + async def consume_one_result(): + generator = manager.generate("multi-turn prompt", None, None, None) + try: + assert await generator.__anext__() == "result" + finally: + await generator.aclose() + + if is_rejected: + with pytest.raises(ServerBusyError): + asyncio.run(consume_one_result()) + assert manager.running_request_count == running_request_count + else: + asyncio.run(consume_one_result()) + assert manager.running_request_count == running_request_count + + def test_pd_master_restores_request_count_when_preload_fails(): class FailingMultimodalParams: async def verify_and_preload(self, request): raise RuntimeError("preload failed") manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs(disable_pd_master_decode_capacity_limit=True) manager.running_request_count = 0 async def consume_generate(): @@ -282,6 +394,7 @@ async def consume_generate(): def test_pd_master_request_count_covers_async_generator_lifecycle(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs(disable_pd_master_decode_capacity_limit=True) manager.running_request_count = 0 inner_generator_closed = False From d9c3d48dafc0946b5fe00f5226663f35d6e4ddea Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 31 Aug 2026 15:50:29 +0800 Subject: [PATCH 2/9] perf(pd): reuse cache match during admission and selection --- .../pd_selector/cache_aware.py | 33 ++++++++++-- unit_tests/server/test_pd_cache_aware.py | 54 +++++++++++++++++++ 2 files changed, 84 insertions(+), 3 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index 08d7a6e90..b143d41be 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -19,13 +19,14 @@ from __future__ import annotations +from contextvars import ContextVar from dataclasses import dataclass from typing import List, Optional from lightllm.server.pd_io_struct import PD_Client_Obj from lightllm.utils.log_utils import init_logger -from .prompt_cache_tree import PromptCacheTree +from .prompt_cache_tree import PromptCacheMatchResult, PromptCacheTree logger = init_logger(__name__) @@ -56,6 +57,18 @@ class CacheAwareConfig: recursion_limit: int = 4000 +@dataclass(frozen=True, slots=True) +class _PromptCacheMatchContext: + policy: "CacheAwarePolicy" + request_text: str + match_result: PromptCacheMatchResult + + +_prompt_cache_match_context: ContextVar[Optional[_PromptCacheMatchContext]] = ContextVar( + "prompt_cache_match_context", default=None +) + + class BalanceRelThresholdController: """根据最近请求的 prompt cache 命中率动态调整负载均衡阈值。""" @@ -157,20 +170,34 @@ def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: self.balance_rel_threshold_controller.update_config(self.config) def estimate_cache_hit_rate(self, workers: List[PD_Client_Obj], request_text: str) -> float: - """Estimate reusable prompt cache on currently connected prefill workers.""" + """Estimate reusable cache and retain the match in this async request context.""" if not workers or not request_text: + _prompt_cache_match_context.set(None) return 0.0 result = self.prompt_cache_tree.prefix_match(request_text) + _prompt_cache_match_context.set( + _PromptCacheMatchContext(policy=self, request_text=request_text, match_result=result) + ) if result.prefill_node is None or not any(worker.client_ip_port == result.prefill_node for worker in workers): return 0.0 if result.input_char_count == 0: return 0.0 return min(max(result.matched_char_count / result.input_char_count, 0.0), 1.0) + def _match_prompt_cache(self, request_text: str) -> PromptCacheMatchResult: + match_context = _prompt_cache_match_context.get() + # Admission and selection share the same prompt object. Identity avoids comparing a potentially long string, + # while ContextVar safely propagates the snapshot to every n-choice child task. + if match_context is not None and match_context.policy is self and match_context.request_text is request_text: + _prompt_cache_match_context.set(None) + return match_context.match_result + _prompt_cache_match_context.set(None) + return self.prompt_cache_tree.prefix_match(request_text) + def _get_cache_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Optional[PD_Client_Obj]: """在指定候选节点中返回达到匹配阈值的 cache 节点。""" - result = self.prompt_cache_tree.prefix_match(request_text) + result = self._match_prompt_cache(request_text) match_rate = 0.0 if result.input_char_count == 0 else result.matched_char_count / result.input_char_count logger.info( diff --git a/unit_tests/server/test_pd_cache_aware.py b/unit_tests/server/test_pd_cache_aware.py index 45fe09173..1e141110d 100644 --- a/unit_tests/server/test_pd_cache_aware.py +++ b/unit_tests/server/test_pd_cache_aware.py @@ -1,3 +1,4 @@ +import asyncio from types import SimpleNamespace import pytest @@ -84,6 +85,59 @@ def test_cache_aware_estimates_hit_rate_only_for_connected_worker(): assert policy.estimate_cache_hit_rate([_worker("10.0.0.2:8000")], prompt) == 0.0 +def test_cache_aware_reuses_admission_match_during_worker_selection(monkeypatch): + policy = CacheAwarePolicy() + cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=110, dispatched_req_num=2) + least_loaded_worker = _worker("10.0.0.2:8000", dispatched_prompt_chars=100, dispatched_req_num=2) + prompt = "shared prefix " * 100 + policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port) + + prefix_match = policy.prompt_cache_tree.prefix_match + match_call_count = 0 + + def counting_prefix_match(text): + nonlocal match_call_count + match_call_count += 1 + return prefix_match(text) + + monkeypatch.setattr(policy.prompt_cache_tree, "prefix_match", counting_prefix_match) + + async def select_worker(): + return policy.select_worker([cache_worker, least_loaded_worker], prompt) + + async def estimate_then_select_in_child_task(): + policy.estimate_cache_hit_rate([cache_worker, least_loaded_worker], prompt) + return await asyncio.gather(*(asyncio.create_task(select_worker()) for _ in range(2))) + + selected_workers = asyncio.run(estimate_then_select_in_child_task()) + assert selected_workers == [cache_worker, cache_worker] + assert match_call_count == 1 + + policy.select_worker([cache_worker, least_loaded_worker], prompt + " new turn") + assert match_call_count == 2 + + +def test_cache_aware_keeps_reused_matches_isolated_between_requests(): + policy = CacheAwarePolicy() + workers = [ + _worker("10.0.0.1:8000", dispatched_req_num=1), + _worker("10.0.0.2:8000", dispatched_req_num=1), + ] + prompts = ["a" * 1024, "b" * 1024] + for worker, prompt in zip(workers, prompts): + policy.prompt_cache_tree.insert(prompt, worker.client_ip_port) + + async def estimate_then_select(prompt): + policy.estimate_cache_hit_rate(workers, prompt) + await asyncio.sleep(0) + return policy.select_worker(workers, prompt) + + async def run_concurrent_requests(): + return await asyncio.gather(*(estimate_then_select(prompt) for prompt in prompts)) + + assert asyncio.run(run_concurrent_requests()) == workers + + def test_cache_aware_keeps_cache_worker_when_inflight_load_is_balanced(): policy = CacheAwarePolicy() cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=110, dispatched_req_num=2) From b915415160fabedcbebf1edf37ddd6b96fa3114e Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 31 Aug 2026 16:59:20 +0800 Subject: [PATCH 3/9] refactor(pd): make waiting admission self-tuning --- lightllm/server/api_cli.py | 19 ---- lightllm/server/core/objs/start_args_type.py | 2 - .../httpserver_for_pd_master/manager.py | 36 +++---- unit_tests/server/test_pd_master_mode.py | 99 +++++-------------- 4 files changed, 37 insertions(+), 119 deletions(-) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 447fa2e5e..60b5fad4e 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -73,25 +73,6 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: action="store_true", help="Disable PD master admission control based on the total capacity of registered decode nodes.", ) - parser.add_argument( - "--pd_master_decode_waiting_queue_ratio", - type=float, - default=0.5, - help=( - "Extra PD master waiting queue size as a ratio of the total capacity of registered decode nodes. " - "For example, 0.5 allows an additional waiting queue equal to 50%% of decode capacity. Default: 0.5." - ), - ) - parser.add_argument( - "--pd_master_cache_aware_queue_reserved_ratio", - type=float, - default=0.5, - help=( - "Fraction of the PD master waiting queue reserved for requests with reusable prompt cache. " - "Cache-aware admission gradually unlocks this reserved capacity based on the estimated cache hit rate. " - "Set to 0 to disable cache-aware reservation. Default: 0.5." - ), - ) parser.add_argument( "--pd_trans_mode", type=str, diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 0f14e53d8..a9aef608b 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -23,8 +23,6 @@ class StartArgs: pd_master_port: int = field(default=1212) pd_master_mode: str = field(default="elastic") disable_pd_master_decode_capacity_limit: bool = field(default=False) - pd_master_decode_waiting_queue_ratio: float = field(default=0.5) - pd_master_cache_aware_queue_reserved_ratio: float = field(default=0.5) pd_trans_mode: str = field(default="nccl", metadata={"choices": ["nccl", "nixl"]}) config_server_host: str = field(default=None) config_server_port: int = field(default=None) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 4ac9ffc3f..58dba3423 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -37,13 +37,6 @@ def __init__( args: StartArgs, ): self.args = args - if ( - not math.isfinite(args.pd_master_decode_waiting_queue_ratio) - or args.pd_master_decode_waiting_queue_ratio < 0 - ): - raise ValueError("pd_master_decode_waiting_queue_ratio must be a non-negative finite number") - if not 0 <= args.pd_master_cache_aware_queue_reserved_ratio <= 1: - raise ValueError("pd_master_cache_aware_queue_reserved_ratio must be between 0 and 1") self.max_req_total_len = args.max_req_total_len assert self.max_req_total_len is not None self.metric_client = MetricClient(get_shm_port_args().metric_port) @@ -139,26 +132,25 @@ async def generate( ): if not self.args.disable_pd_master_decode_capacity_limit: decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.pd_manager.decode_nodes) - waiting_queue_capacity = math.ceil(decode_capacity * self.args.pd_master_decode_waiting_queue_ratio) - hard_admission_limit = decode_capacity + waiting_queue_capacity + # Every request gets half a decode wave of buffering. Reusable prompt cache progressively unlocks + # the other half, while two full decode waves remain the hard upper bound. + general_waiting_capacity = (decode_capacity + 1) // 2 + cache_waiting_capacity = decode_capacity // 2 + hard_admission_limit = decode_capacity + general_waiting_capacity + cache_waiting_capacity if self.running_request_count >= hard_admission_limit: raise ServerBusyError() - reserved_ratio = self.args.pd_master_cache_aware_queue_reserved_ratio - general_waiting_capacity = math.ceil(waiting_queue_capacity * (1.0 - reserved_ratio)) general_admission_limit = decode_capacity + general_waiting_capacity if self.running_request_count >= general_admission_limit: - estimate_cache_hit_rate = getattr(self.pd_manager.selector, "estimate_prompt_cache_hit_rate", None) - estimated_cache_hit_rate = estimate_cache_hit_rate(prompt) if estimate_cache_hit_rate else None - if estimated_cache_hit_rate is not None: - if not math.isfinite(estimated_cache_hit_rate): - estimated_cache_hit_rate = 0.0 - estimated_cache_hit_rate = min(max(estimated_cache_hit_rate, 0.0), 1.0) - request_waiting_capacity = math.ceil( - waiting_queue_capacity * (1.0 - reserved_ratio * (1.0 - estimated_cache_hit_rate)) - ) - if self.running_request_count >= decode_capacity + request_waiting_capacity: - raise ServerBusyError() + estimated_cache_hit_rate = self.pd_manager.selector.estimate_prompt_cache_hit_rate(prompt) + if estimated_cache_hit_rate is None or not math.isfinite(estimated_cache_hit_rate): + estimated_cache_hit_rate = 0.0 + estimated_cache_hit_rate = min(max(estimated_cache_hit_rate, 0.0), 1.0) + cache_aware_admission_limit = general_admission_limit + math.ceil( + cache_waiting_capacity * estimated_cache_hit_rate + ) + if self.running_request_count >= cache_aware_admission_limit: + raise ServerBusyError() was_idle = self.running_request_count == 0 self.running_request_count += 1 diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index fed4563f0..b7e06b140 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -5,7 +5,6 @@ import pytest from easydict import EasyDict -from lightllm.server.api_cli import make_argument_parser from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager from lightllm.utils.error_utils import ServerBusyError @@ -87,22 +86,6 @@ def test_pd_master_models_endpoint_has_created_timestamp(monkeypatch): assert response.data[0].created == 1234 -def test_pd_master_decode_waiting_queue_ratio_cli_default_and_override(): - default_args = make_argument_parser().parse_args([]) - assert default_args.pd_master_decode_waiting_queue_ratio == 0.5 - assert default_args.pd_master_cache_aware_queue_reserved_ratio == 0.5 - configured_args = make_argument_parser().parse_args( - [ - "--pd_master_decode_waiting_queue_ratio", - "0.25", - "--pd_master_cache_aware_queue_reserved_ratio", - "0.75", - ] - ) - assert configured_args.pd_master_decode_waiting_queue_ratio == 0.25 - assert configured_args.pd_master_cache_aware_queue_reserved_ratio == 0.75 - - def test_elastic_pd_nodes_are_ready_with_at_least_one_node_of_each_role(): manager = PDManager(StartArgs(pd_master_mode="elastic")) assert manager.is_pd_nodes_ready() is False @@ -281,73 +264,33 @@ def test_pd_master_inference_health_matches_normal_node_semantics(monkeypatch): @pytest.mark.parametrize( - ("waiting_queue_ratio", "running_request_count", "is_rejected"), + ("decode_capacity", "estimated_cache_hit_rate", "running_request_count", "is_rejected"), [ - (0.5, 11, False), - (0.5, 12, True), - (0.25, 9, False), - (0.25, 10, True), - (0.0, 7, False), - (0.0, 8, True), + (8, None, 11, False), + (8, None, 12, True), + (3, None, 4, False), + (3, None, 5, True), + (8, 0.0, 12, True), + (8, 0.5, 13, False), + (8, 0.5, 14, True), + (8, 1.0, 15, False), + (8, 1.0, 16, True), ], ) -def test_pd_master_admission_limit_includes_configurable_waiting_queue( - waiting_queue_ratio, running_request_count, is_rejected +def test_pd_master_admission_adapts_to_capacity_and_cache_hit_rate( + decode_capacity, estimated_cache_hit_rate, running_request_count, is_rejected ): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) - manager.args = StartArgs(pd_master_decode_waiting_queue_ratio=waiting_queue_ratio) - manager.pd_manager = SimpleNamespace( - decode_nodes=[ - SimpleNamespace(start_args={"running_max_req_size": 3}), - SimpleNamespace(start_args={"running_max_req_size": 5}), - ], - selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: None), - ) - manager.running_request_count = running_request_count - manager.latest_success_infer_time = 0 - - async def fake_generate(*_args): - yield "result" - - manager._generate = fake_generate - - async def consume_one_result(): - generator = manager.generate(None, None, None, None) - try: - assert await generator.__anext__() == "result" - finally: - await generator.aclose() - - if is_rejected: - with pytest.raises(ServerBusyError): - asyncio.run(consume_one_result()) - assert manager.running_request_count == running_request_count - else: - asyncio.run(consume_one_result()) - assert manager.running_request_count == running_request_count + manager.args = StartArgs() + estimate_calls = [] + def estimate_prompt_cache_hit_rate(prompt): + estimate_calls.append(prompt) + return estimated_cache_hit_rate -@pytest.mark.parametrize( - ("estimated_cache_hit_rate", "running_request_count", "is_rejected"), - [ - (0.0, 9, False), - (0.0, 10, True), - (0.5, 10, False), - (0.5, 11, True), - (1.0, 11, False), - (1.0, 12, True), - (None, 11, False), - (None, 12, True), - ], -) -def test_pd_master_admission_reserves_queue_for_cache_friendly_requests( - estimated_cache_hit_rate, running_request_count, is_rejected -): - manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) - manager.args = StartArgs() manager.pd_manager = SimpleNamespace( - decode_nodes=[SimpleNamespace(start_args={"running_max_req_size": 8})], - selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: estimated_cache_hit_rate), + decode_nodes=[SimpleNamespace(start_args={"running_max_req_size": decode_capacity})], + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=estimate_prompt_cache_hit_rate), ) manager.running_request_count = running_request_count manager.latest_success_infer_time = 0 @@ -372,6 +315,10 @@ async def consume_one_result(): asyncio.run(consume_one_result()) assert manager.running_request_count == running_request_count + general_admission_limit = decode_capacity + (decode_capacity + 1) // 2 + should_estimate_cache = general_admission_limit <= running_request_count < 2 * decode_capacity + assert len(estimate_calls) == int(should_estimate_cache) + def test_pd_master_restores_request_count_when_preload_fails(): class FailingMultimodalParams: From c31848a9818052af6e195bdb012b5bebbb364e42 Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 31 Aug 2026 20:54:04 +0800 Subject: [PATCH 4/9] docs: translate PD admission comments to Chinese --- lightllm/server/httpserver_for_pd_master/manager.py | 4 ++-- .../httpserver_for_pd_master/pd_selector/cache_aware.py | 6 +++--- .../httpserver_for_pd_master/pd_selector/pd_selector.py | 2 +- 3 files changed, 6 insertions(+), 6 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 58dba3423..6931dd0d3 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -132,8 +132,8 @@ async def generate( ): if not self.args.disable_pd_master_decode_capacity_limit: decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.pd_manager.decode_nodes) - # Every request gets half a decode wave of buffering. Reusable prompt cache progressively unlocks - # the other half, while two full decode waves remain the hard upper bound. + # 每个请求默认获得半个解码波次的缓冲容量;可复用的提示词缓存会逐步开放另外半个波次, + # 同时以两个完整解码波次作为硬上限。 general_waiting_capacity = (decode_capacity + 1) // 2 cache_waiting_capacity = decode_capacity // 2 hard_admission_limit = decode_capacity + general_waiting_capacity + cache_waiting_capacity diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index b143d41be..4a5357c93 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -170,7 +170,7 @@ def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: self.balance_rel_threshold_controller.update_config(self.config) def estimate_cache_hit_rate(self, workers: List[PD_Client_Obj], request_text: str) -> float: - """Estimate reusable cache and retain the match in this async request context.""" + """估算可复用缓存的命中率,并在当前异步请求上下文中保留匹配结果。""" if not workers or not request_text: _prompt_cache_match_context.set(None) return 0.0 @@ -187,8 +187,8 @@ def estimate_cache_hit_rate(self, workers: List[PD_Client_Obj], request_text: st def _match_prompt_cache(self, request_text: str) -> PromptCacheMatchResult: match_context = _prompt_cache_match_context.get() - # Admission and selection share the same prompt object. Identity avoids comparing a potentially long string, - # while ContextVar safely propagates the snapshot to every n-choice child task. + # 准入和选点共享同一个提示词对象。通过对象身份判断可以避免比较可能很长的字符串, + # ContextVar 则会将快照安全地传递给每个 n-choice 子任务。 if match_context is not None and match_context.policy is self and match_context.request_text is request_text: _prompt_cache_match_context.set(None) return match_context.match_result diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py index a8c3a6e04..7bf844341 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py @@ -31,7 +31,7 @@ def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: return def estimate_prompt_cache_hit_rate(self, prompt: Union[str, List[int]]) -> Optional[float]: - """Return None when this selector cannot estimate reusable prompt cache.""" + """当选择器无法估算可复用的提示词缓存时返回 None。""" return None From c72f1590686ef936e9bb95b6594f19c8025ff019 Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 31 Aug 2026 22:35:27 +0800 Subject: [PATCH 5/9] feat(pd): add adaptive cache-aware admission control --- lightllm/server/api_cli.py | 2 +- lightllm/server/api_http_pd.py | 2 + lightllm/server/httpserver/pd_loop.py | 116 +++- .../httpserver_for_pd_master/admission.py | 520 ++++++++++++++++++ .../httpserver_for_pd_master/manager.py | 197 ++++++- lightllm/server/metrics/metrics.py | 6 + lightllm/server/pd_io_struct.py | 11 + unit_tests/server/test_pd_admission.py | 333 +++++++++++ unit_tests/server/test_pd_master_mode.py | 247 +++++++-- 9 files changed, 1348 insertions(+), 86 deletions(-) create mode 100644 lightllm/server/httpserver_for_pd_master/admission.py create mode 100644 unit_tests/server/test_pd_admission.py diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 60b5fad4e..16739babf 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -71,7 +71,7 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: parser.add_argument( "--disable_pd_master_decode_capacity_limit", action="store_true", - help="Disable PD master admission control based on the total capacity of registered decode nodes.", + help="Disable the PD master capacity and cache-aware admission queue.", ) parser.add_argument( "--pd_trans_mode", diff --git a/lightllm/server/api_http_pd.py b/lightllm/server/api_http_pd.py index 1d8b2112f..26c27e6fe 100644 --- a/lightllm/server/api_http_pd.py +++ b/lightllm/server/api_http_pd.py @@ -41,6 +41,8 @@ async def register_and_keep_alive(websocket: WebSocket): data = await asyncio.wait_for(websocket.receive_bytes(), timeout=heartbeat_timeout_seconds) obj = pickle.loads(data) if isinstance(obj, tuple) and obj and obj[0] == ObjType.HEARTBEAT: + load_info = obj[1] if len(obj) > 1 else None + g_objs.httpserver_manager.update_node_load_info(load_info) continue await g_objs.httpserver_manager.put_to_handle_queue(obj) diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index e66540cb5..61656459a 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -9,22 +9,86 @@ import os import signal import sys +import time from typing import Dict, Optional, Union, List from websockets import ClientConnection from lightllm.server.pd_io_struct import NodeRole, ObjType from lightllm.server.httpserver.async_queue import AsyncQueue from lightllm.utils.net_utils import get_hostname_ip from lightllm.utils.log_utils import init_logger -from lightllm.utils.envs_utils import get_lightllm_websocket_max_message_size +from lightllm.utils.envs_utils import get_lightllm_websocket_max_message_size, get_unique_server_name from lightllm.server.httpserver.manager import HttpServerManager from ..pd_io_struct import PD_Master_Obj from lightllm.server.core.objs import StartArgs from lightllm.server.core.objs import SamplingParams from lightllm.utils.error_utils import PDPrefillNodeStopGenToken from lightllm.utils.shm_port_args import get_shm_port_args +from lightllm.server.router.dynamic_prompt.radix_cache import RadixCacheReadOnlyClient logger = init_logger(__name__) +_radix_cache_client = None +_radix_cache_client_key = None + + +def _update_pd_master_membership(manager: HttpServerManager, pd_master_ids) -> None: + pd_master_ids = tuple(sorted(pd_master_ids)) + if getattr(manager, "pd_master_ids", ()) == pd_master_ids: + return + manager.pd_master_ids = pd_master_ids + manager.pd_master_capacity_epoch = max( + getattr(manager, "pd_master_capacity_epoch", 0) + 1, + time.time_ns(), + ) + membership_changed = getattr(manager, "pd_master_membership_changed", None) + if membership_changed is None: + membership_changed = manager.pd_master_membership_changed = asyncio.Event() + membership_changed.set() + + +def _allocate_capacity_share(total_capacity: int, pd_master_ids, pd_master_node_id: int) -> int: + """把节点容量确定性地切成互不重叠的 PD Master 租约池。""" + pd_master_ids = tuple(sorted(pd_master_ids)) + if not pd_master_ids or pd_master_node_id not in pd_master_ids: + return 0 + base, remainder = divmod(max(0, total_capacity), len(pd_master_ids)) + return base + int(pd_master_ids.index(pd_master_node_id) < remainder) + + +def _get_radix_cache_info(): + global _radix_cache_client, _radix_cache_client_key + + from lightllm.server.api_http import g_objs + + args = g_objs.args + if args.disable_dynamic_prompt_cache: + return 0, 0, 0 + + max_total_token_num = g_objs.httpserver_manager.shm_max_total_token_num.get_value() + if max_total_token_num <= 0: + return 0, 0, 0 + + node_world_size = args.tp // args.nnodes + dp_world_size = args.tp // args.dp + client_key = (get_unique_server_name(), max_total_token_num, node_world_size, dp_world_size) + try: + if _radix_cache_client is None or _radix_cache_client_key != client_key: + _radix_cache_client = RadixCacheReadOnlyClient( + get_unique_server_name(), + max_total_token_num, + node_world_size=node_world_size, + dp_world_size=dp_world_size, + ) + _radix_cache_client_key = client_key + + dp_size_in_node = max(1, args.dp // args.nnodes) + total_tokens = sum(_radix_cache_client.get_tree_total_tokens_num(i) for i in range(dp_size_in_node)) + refed_tokens = sum(_radix_cache_client.get_refed_tokens_num(i) for i in range(dp_size_in_node)) + return int(total_tokens), int(refed_tokens), int(max_total_token_num * dp_size_in_node) + except Exception as exc: + logger.debug(f"read radix cache load failed: {str(exc)}") + return 0, 0, 0 + async def timer_log(manager: HttpServerManager): while True: @@ -56,7 +120,8 @@ async def pd_handle_loop(manager: HttpServerManager): logger.info(f"get pd_master_objs {id_to_pd_master_obj}") if id_to_pd_master_obj is not None: - for node_id, pd_master_obj in id_to_handle_task.items(): + _update_pd_master_membership(manager, id_to_pd_master_obj) + for node_id, pd_master_obj in list(id_to_handle_task.items()): if node_id not in id_to_pd_master_obj: id_to_handle_task[node_id].cancel() id_to_handle_task.pop(node_id, None) @@ -106,14 +171,24 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O "client_ip_port": f"{manager.host_ip}:{get_shm_port_args().port}", "mode": manager.pd_mode.value, "start_args": args_dict, + "capacity_share": _allocate_capacity_share( + manager.args.running_max_req_size, + manager.pd_master_ids, + pd_master_obj.node_id, + ), + "capacity_epoch": manager.pd_master_capacity_epoch, } await websocket.send(json.dumps(regist_json)) logger.info(f"Sent registration JSON: {regist_json}") # 转发任务 - forwarding_tokens_task = asyncio.create_task(_up_tokens_to_pd_master(forwarding_queue, websocket)) - heartbeat_task = asyncio.create_task(_send_heartbeat_to_pd_master(websocket)) + forwarding_tokens_task = asyncio.create_task( + _up_tokens_to_pd_master(forwarding_queue, websocket, pd_master_obj.node_id) + ) + heartbeat_task = asyncio.create_task( + _send_heartbeat_to_pd_master(manager, websocket, pd_master_obj.node_id) + ) group_req_id_to_event: Dict[int, asyncio.Event] = weakref.WeakValueDictionary() # 接收 pd master 发来的请求,并推理后,将生成的token转发回pd master。 @@ -264,24 +339,38 @@ async def _pd_process_generate( # 转发token的task -async def _up_tokens_to_pd_master(forwarding_queue: AsyncQueue, websocket: ClientConnection): +async def _up_tokens_to_pd_master( + forwarding_queue: AsyncQueue, + websocket: ClientConnection, + pd_master_node_id: int, +): while True: handle_list = await forwarding_queue.wait_to_get_all_data() if handle_list: - load_info: dict = _get_load_info() + load_info: dict = _get_load_info(pd_master_node_id) await websocket.send(pickle.dumps((ObjType.TOKEN_PACKS, handle_list, load_info))) -async def _send_heartbeat_to_pd_master(websocket: ClientConnection): +async def _send_heartbeat_to_pd_master( + manager: HttpServerManager, + websocket: ClientConnection, + pd_master_node_id: int, +): heartbeat_interval_seconds = 15 + membership_changed = manager.pd_master_membership_changed while True: - await websocket.send(pickle.dumps((ObjType.HEARTBEAT,))) - await asyncio.sleep(heartbeat_interval_seconds) + membership_changed.clear() + await websocket.send(pickle.dumps((ObjType.HEARTBEAT, _get_load_info(pd_master_node_id)))) + try: + # Master 集合变化时立即重报份额,缩短新旧容量租约并存的窗口。 + await asyncio.wait_for(membership_changed.wait(), timeout=heartbeat_interval_seconds) + except asyncio.TimeoutError: + pass # 获取节点负载信息 -def _get_load_info() -> dict: +def _get_load_info(pd_master_node_id: int) -> dict: from lightllm.server.api_http import g_objs @@ -295,8 +384,15 @@ def _get_load_info() -> dict: float(g_objs.shared_token_load.get_dynamic_max_load(dp_index)) for dp_index in range(dp_size_in_node) ] mean_node_load = sum(current_load) / len(current_load) + radix_cache_total_tokens, radix_cache_refed_tokens, radix_cache_capacity_tokens = _get_radix_cache_info() + pd_master_ids = getattr(g_objs.httpserver_manager, "pd_master_ids", (pd_master_node_id,)) load_info = { "total_token_usage_rate": mean_node_load, "client_ip_port": f"{g_objs.httpserver_manager.host_ip}:{get_shm_port_args().port}", + "capacity_share": _allocate_capacity_share(args.running_max_req_size, pd_master_ids, pd_master_node_id), + "capacity_epoch": getattr(g_objs.httpserver_manager, "pd_master_capacity_epoch", 0), + "radix_cache_total_tokens": radix_cache_total_tokens, + "radix_cache_refed_tokens": radix_cache_refed_tokens, + "radix_cache_capacity_tokens": radix_cache_capacity_tokens, } return load_info diff --git a/lightllm/server/httpserver_for_pd_master/admission.py b/lightllm/server/httpserver_for_pd_master/admission.py new file mode 100644 index 000000000..6cda68274 --- /dev/null +++ b/lightllm/server/httpserver_for_pd_master/admission.py @@ -0,0 +1,520 @@ +from __future__ import annotations + +import asyncio +import time +from collections import OrderedDict, deque +from dataclasses import dataclass, replace +from enum import IntEnum +from typing import Callable, Deque, Dict, Optional + +from lightllm.utils.error_utils import ServerBusyError + + +class AdmissionPriority(IntEnum): + """请求的业务优先级;数值越大,获得服务的权重越高。""" + + COLD = 0 + PROBABLE_CACHE_HIT = 1 + CONTINUATION = 2 + + +@dataclass(frozen=True, slots=True) +class AdmissionPolicy: + """PD Master 内部准入策略。 + + 这些值表达产品层面的等待与公平策略,因此集中在一个对象中,不散落到 + 调度控制流里。容量本身由已注册的 Decode 节点动态提供。 + """ + + continuation_weight: int = 8 + probable_cache_hit_weight: int = 3 + cold_weight: int = 1 + continuation_max_wait_seconds: float = 30.0 + probable_cache_hit_max_wait_seconds: float = 15.0 + cold_max_wait_seconds: float = 5.0 + waiting_decode_waves: int = 1 + probable_cache_hit_threshold: float = 0.5 + active_session_ttl_seconds: float = 30 * 60.0 + max_tracked_sessions: int = 100_000 + + def __post_init__(self) -> None: + if ( + min( + self.continuation_weight, + self.probable_cache_hit_weight, + self.cold_weight, + ) + < 1 + ): + raise ValueError("admission weights must be positive") + if ( + min( + self.continuation_max_wait_seconds, + self.probable_cache_hit_max_wait_seconds, + self.cold_max_wait_seconds, + ) + <= 0 + ): + raise ValueError("admission wait timeouts must be positive") + if self.waiting_decode_waves < 1: + raise ValueError("waiting_decode_waves must be positive") + if not 0.0 <= self.probable_cache_hit_threshold <= 1.0: + raise ValueError("probable_cache_hit_threshold must be between zero and one") + if self.active_session_ttl_seconds <= 0: + raise ValueError("active_session_ttl_seconds must be positive") + if self.max_tracked_sessions < 1: + raise ValueError("max_tracked_sessions must be positive") + + def weight(self, priority: AdmissionPriority) -> int: + if priority == AdmissionPriority.CONTINUATION: + return self.continuation_weight + if priority == AdmissionPriority.PROBABLE_CACHE_HIT: + return self.probable_cache_hit_weight + return self.cold_weight + + def max_wait_seconds(self, priority: AdmissionPriority) -> float: + if priority == AdmissionPriority.CONTINUATION: + return self.continuation_max_wait_seconds + if priority == AdmissionPriority.PROBABLE_CACHE_HIT: + return self.probable_cache_hit_max_wait_seconds + return self.cold_max_wait_seconds + + +@dataclass(frozen=True, slots=True) +class AdmissionRequest: + session_key: Optional[str] + priority: AdmissionPriority + decode_slots: int = 1 + estimated_uncached_work: int = 0 + + def __post_init__(self) -> None: + if self.decode_slots < 1: + raise ValueError("decode_slots must be positive") + if self.estimated_uncached_work < 0: + raise ValueError("estimated_uncached_work must be non-negative") + + +@dataclass(frozen=True, slots=True) +class CacheCapacitySnapshot: + """当前 PD Master 可使用的 Prefill Radix cache 份额。""" + + total_tokens: int + capacity_tokens: int + + def __post_init__(self) -> None: + if self.total_tokens < 0 or self.capacity_tokens < 0: + raise ValueError("cache token counts must be non-negative") + + @property + def free_tokens(self) -> int: + return max(0, self.capacity_tokens - self.total_tokens) + + +class SessionTracker: + """只把服务端已经成功观察过的 Session 视为连续会话。""" + + def __init__( + self, + ttl_seconds: float, + max_sessions: int, + clock: Callable[[], float] = time.monotonic, + ) -> None: + self.ttl_seconds = ttl_seconds + self.max_sessions = max_sessions + self._clock = clock + self._last_success: OrderedDict[str, float] = OrderedDict() + + def is_continuation(self, session_key: Optional[str]) -> bool: + if not session_key: + return False + now = self._clock() + last_success = self._last_success.get(session_key) + if last_success is None: + return False + if now - last_success > self.ttl_seconds: + self._last_success.pop(session_key, None) + return False + self._last_success.move_to_end(session_key) + return True + + def mark_success(self, session_key: Optional[str]) -> None: + if not session_key: + return + self._last_success[session_key] = self._clock() + self._last_success.move_to_end(session_key) + while len(self._last_success) > self.max_sessions: + self._last_success.popitem(last=False) + + +@dataclass(slots=True) +class _WaitingRequest: + sequence_id: int + request: AdmissionRequest + enqueue_time: float + future: asyncio.Future + + +class AdmissionLease: + """一次已经获得的 Decode 容量租约。""" + + def __init__( + self, + controller: "PDAdmissionController", + request: AdmissionRequest, + waited_seconds: float, + ) -> None: + self._controller = controller + self.request = request + self.waited_seconds = waited_seconds + self._released = False + + async def __aenter__(self) -> "AdmissionLease": + return self + + async def __aexit__(self, _exc_type, _exc, _traceback) -> None: + self.release() + + def release(self) -> None: + if self._released: + return + self._released = True + self._controller._release(self) + + +class PDAdmissionController: + """在请求派发到 P/D 节点之前提供有界、可取消的公平等待队列。""" + + _PRIORITY_ORDER = ( + AdmissionPriority.CONTINUATION, + AdmissionPriority.PROBABLE_CACHE_HIT, + AdmissionPriority.COLD, + ) + + def __init__( + self, + decode_capacity_provider: Callable[[], int], + cache_capacity_provider: Optional[Callable[[], Optional[CacheCapacitySnapshot]]] = None, + policy: Optional[AdmissionPolicy] = None, + clock: Callable[[], float] = time.monotonic, + state_change_callback: Optional[Callable[["PDAdmissionController"], None]] = None, + ) -> None: + self.policy = policy or AdmissionPolicy() + self._decode_capacity_provider = decode_capacity_provider + self._cache_capacity_provider = cache_capacity_provider + self._clock = clock + self._state_change_callback = state_change_callback + self._active_slots = 0 + self._active_cold_slots = 0 + self._average_cold_uncached_tokens: Optional[float] = None + self._active_sessions = set() + self._queues: Dict[AdmissionPriority, Deque[_WaitingRequest]] = { + priority: deque() for priority in self._PRIORITY_ORDER + } + self._session_queues: Dict[str, Deque[_WaitingRequest]] = {} + self._queued_slots = 0 + self._sequence_id = 0 + self._schedule = self._build_schedule() + self._schedule_index = 0 + + @property + def active_slots(self) -> int: + return self._active_slots + + @property + def active_cold_slots(self) -> int: + return self._active_cold_slots + + @property + def cold_capacity(self) -> int: + """返回在当前缓存余量下允许并发的冷请求槽位数。""" + decode_capacity = self._capacity() + if ( + decode_capacity <= 0 + or self._cache_capacity_provider is None + or self._average_cold_uncached_tokens is None + or self._average_cold_uncached_tokens <= 0 + ): + return decode_capacity + + snapshot = self._cache_capacity_provider() + if snapshot is None or snapshot.capacity_tokens <= 0: + return decode_capacity + + # 剩余缓存能容纳几个“平均冷请求”,就开放几个冷槽位;至少保留一个 + # 探索槽位,使系统在缓存已满时仍能接纳新会话并持续获得反馈。 + requests_fitting_in_cache = int(snapshot.free_tokens / self._average_cold_uncached_tokens) + return min(decode_capacity, max(1, requests_fitting_in_cache)) + + @property + def queued_slots(self) -> int: + return self._queued_slots + + @property + def queued_request_count(self) -> int: + return sum(len(queue) for queue in self._queues.values()) + + def _capacity(self) -> int: + return max(0, int(self._decode_capacity_provider())) + + def _build_schedule(self) -> tuple[AdmissionPriority, ...]: + remaining = {priority: self.policy.weight(priority) for priority in self._PRIORITY_ORDER} + schedule = [] + while any(remaining.values()): + for priority in self._PRIORITY_ORDER: + if remaining[priority] > 0: + schedule.append(priority) + remaining[priority] -= 1 + return tuple(schedule) + + async def acquire(self, request: AdmissionRequest) -> AdmissionLease: + capacity = self._capacity() + if capacity <= 0 or request.decode_slots > capacity: + raise ServerBusyError("PD decode capacity is unavailable") + + if self.queued_request_count == 0 and self._can_activate(request): + lease = self._activate(request, waited_seconds=0.0) + self._notify_state_change() + return lease + + loop = asyncio.get_running_loop() + waiter = _WaitingRequest( + sequence_id=self._sequence_id, + request=request, + enqueue_time=self._clock(), + future=loop.create_future(), + ) + self._sequence_id += 1 + + if not self._make_queue_room(waiter): + raise ServerBusyError("PD master admission queue is full") + + self._enqueue(waiter) + self._drain() + + try: + return await asyncio.wait_for( + asyncio.shield(waiter.future), + timeout=self.policy.max_wait_seconds(request.priority), + ) + except asyncio.TimeoutError as exc: + lease = self._cancel_waiter_or_take_lease(waiter) + if lease is not None: + lease.release() + raise ServerBusyError("PD master admission queue wait timed out") from exc + except BaseException: + lease = self._cancel_waiter_or_take_lease(waiter) + if lease is not None: + lease.release() + raise + + def on_capacity_change(self) -> None: + self._drain() + + def record_prefill_result( + self, + request: AdmissionRequest, + prompt_tokens: int, + cached_tokens: int, + ) -> None: + """用冷请求的真实未命中量更新下一轮冷容量。""" + if request.priority != AdmissionPriority.COLD: + return + + prompt_tokens = max(0, int(prompt_tokens)) + cached_tokens = min(prompt_tokens, max(0, int(cached_tokens))) + uncached_tokens = prompt_tokens - cached_tokens + + # 一个 Decode 波次作为自适应窗口:容量越大,单个样本对均值的影响越小。 + sample_window = max(1, self._capacity()) + alpha = 2.0 / (sample_window + 1.0) + if self._average_cold_uncached_tokens is None: + self._average_cold_uncached_tokens = float(uncached_tokens) + else: + self._average_cold_uncached_tokens += alpha * (uncached_tokens - self._average_cold_uncached_tokens) + self._drain() + + def promote_session(self, session_key: Optional[str]) -> None: + """把同一 Session 尚未派发的请求提升为连续会话优先级。""" + if not session_key: + return + session_queue = self._session_queues.get(session_key) + if not session_queue: + return + + for waiter in tuple(session_queue): + old_priority = waiter.request.priority + if old_priority == AdmissionPriority.CONTINUATION: + continue + self._queues[old_priority].remove(waiter) + waiter.request = replace(waiter.request, priority=AdmissionPriority.CONTINUATION) + self._queues[AdmissionPriority.CONTINUATION].append(waiter) + self._drain() + + def _has_cold_capacity(self, request: AdmissionRequest) -> bool: + if request.priority != AdmissionPriority.COLD: + return True + return self._active_cold_slots + request.decode_slots <= self.cold_capacity + + def _can_activate(self, request: AdmissionRequest) -> bool: + if self._active_slots + request.decode_slots > self._capacity(): + return False + if not self._has_cold_capacity(request): + return False + return request.session_key is None or request.session_key not in self._active_sessions + + def _activate(self, request: AdmissionRequest, waited_seconds: float) -> AdmissionLease: + self._active_slots += request.decode_slots + if request.priority == AdmissionPriority.COLD: + self._active_cold_slots += request.decode_slots + if request.session_key is not None: + self._active_sessions.add(request.session_key) + return AdmissionLease(self, request, waited_seconds) + + def _release(self, lease: AdmissionLease) -> None: + request = lease.request + self._active_slots -= request.decode_slots + if request.priority == AdmissionPriority.COLD: + self._active_cold_slots -= request.decode_slots + if self._active_slots < 0: + raise RuntimeError("PD admission active slot count became negative") + if self._active_cold_slots < 0: + raise RuntimeError("PD admission active cold slot count became negative") + if request.session_key is not None: + self._active_sessions.discard(request.session_key) + self._drain() + + def _waiting_capacity(self) -> int: + return self._capacity() * self.policy.waiting_decode_waves + + def _make_queue_room(self, incoming: _WaitingRequest) -> bool: + waiting_capacity = self._waiting_capacity() + if incoming.request.decode_slots > waiting_capacity: + return False + + required_slots = self._queued_slots + incoming.request.decode_slots - waiting_capacity + if required_slots <= 0: + return True + + victims = [] + released_slots = 0 + for priority in reversed(self._PRIORITY_ORDER): + if priority >= incoming.request.priority: + continue + for waiter in reversed(self._queues[priority]): + victims.append(waiter) + released_slots += waiter.request.decode_slots + if released_slots >= required_slots: + break + if released_slots >= required_slots: + break + + if released_slots < required_slots: + return False + + for victim in victims: + self._remove_waiter(victim) + if not victim.future.done(): + victim.future.set_exception(ServerBusyError("Superseded by a higher-priority queued request")) + return True + + def _enqueue(self, waiter: _WaitingRequest) -> None: + self._queues[waiter.request.priority].append(waiter) + self._queued_slots += waiter.request.decode_slots + if waiter.request.session_key is not None: + self._session_queues.setdefault(waiter.request.session_key, deque()).append(waiter) + + def _remove_waiter(self, waiter: _WaitingRequest) -> bool: + try: + self._queues[waiter.request.priority].remove(waiter) + except ValueError: + return False + + self._queued_slots -= waiter.request.decode_slots + session_key = waiter.request.session_key + if session_key is not None: + session_queue = self._session_queues[session_key] + session_queue.remove(waiter) + if not session_queue: + self._session_queues.pop(session_key, None) + return True + + def _cancel_waiter_or_take_lease(self, waiter: _WaitingRequest) -> Optional[AdmissionLease]: + if self._remove_waiter(waiter): + waiter.future.cancel() + self._drain() + return None + if waiter.future.done() and not waiter.future.cancelled(): + try: + result = waiter.future.result() + except BaseException: + return None + if isinstance(result, AdmissionLease): + return result + return None + + def _first_grantable(self, priority: AdmissionPriority) -> Optional[_WaitingRequest]: + candidates = [] + available_decode_slots = self._capacity() - self._active_slots + for waiter in self._queues[priority]: + session_key = waiter.request.session_key + if session_key is not None and session_key in self._active_sessions: + continue + if session_key is not None and self._session_queues[session_key][0] is not waiter: + continue + if priority != AdmissionPriority.COLD: + return waiter + if waiter.request.decode_slots > self.cold_capacity: + continue + if waiter.request.decode_slots <= available_decode_slots and not self._has_cold_capacity(waiter.request): + continue + candidates.append(waiter) + + if not candidates: + return None + # 冷请求内部优先处理预计新增缓存最少的任务,在同等代价下保持 FIFO。 + return min( + candidates, + key=lambda waiter: ( + waiter.request.estimated_uncached_work, + waiter.sequence_id, + ), + ) + + def _select_next(self) -> Optional[_WaitingRequest]: + schedule_size = len(self._schedule) + for offset in range(schedule_size): + index = (self._schedule_index + offset) % schedule_size + waiter = self._first_grantable(self._schedule[index]) + if waiter is not None: + self._schedule_index = (index + 1) % schedule_size + return waiter + return None + + def _drain(self) -> None: + while self._active_slots < self._capacity(): + schedule_index = self._schedule_index + waiter = self._select_next() + if waiter is None: + break + if self._active_slots + waiter.request.decode_slots > self._capacity(): + # 为需要多个 choice slot 的老请求保留逐步释放出来的容量,避免永久饥饿。 + self._schedule_index = schedule_index + break + if not self._has_cold_capacity(waiter.request): + self._schedule_index = schedule_index + break + if not self._remove_waiter(waiter): + continue + lease = self._activate( + waiter.request, + waited_seconds=max(0.0, self._clock() - waiter.enqueue_time), + ) + if waiter.future.done(): + lease.release() + continue + waiter.future.set_result(lease) + self._notify_state_change() + + def _notify_state_change(self) -> None: + if self._state_change_callback is not None: + self._state_change_callback(self) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 6931dd0d3..ea26bd328 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -27,6 +27,14 @@ from lightllm.utils.envs_utils import get_pd_split_max_new_tokens from lightllm.utils.shm_port_args import get_shm_port_args from .pd_selector import create_selector +from .admission import ( + AdmissionPolicy, + AdmissionPriority, + AdmissionRequest, + CacheCapacitySnapshot, + PDAdmissionController, + SessionTracker, +) logger = init_logger(__name__) @@ -44,6 +52,18 @@ def __init__( self.pd_manager = PDManager(args) + self.admission_policy = AdmissionPolicy() + self.session_tracker = SessionTracker( + ttl_seconds=self.admission_policy.active_session_ttl_seconds, + max_sessions=self.admission_policy.max_tracked_sessions, + ) + self.admission_controller = PDAdmissionController( + decode_capacity_provider=self.pd_manager.get_decode_capacity, + cache_capacity_provider=self.pd_manager.get_prefill_cache_capacity, + policy=self.admission_policy, + state_change_callback=self._record_admission_state, + ) + self.req_id_to_out_inf: Dict[int, ReqStatus] = {} self.infos_queues = None # 这个需要延迟初始化,否则使用的loop不对 self.health_timeout = int(os.getenv("HEALTH_TIMEOUT", "200")) @@ -76,10 +96,12 @@ def is_healthy(self): async def register_pd(self, pd_info_json, websocket): self.pd_manager.register_pd(pd_info_json, websocket) + self.admission_controller.on_capacity_change() return async def remove_pd(self, pd_info_json): self.pd_manager.remove_pd(pd_info_json) + self.admission_controller.on_capacity_change() return async def update_req_status(self, upkv_status: PDUpKVStatus): @@ -92,6 +114,25 @@ async def update_req_status(self, upkv_status: PDUpKVStatus): pass return + def update_node_load_info(self, load_info: Optional[dict]) -> None: + self.pd_manager.update_node_load_info(load_info) + # Decode 租约或 Prefill cache 余量变化后都需要重新尝试队列。 + self.admission_controller.on_capacity_change() + + def _record_admission_state(self, controller: PDAdmissionController) -> None: + self.metric_client.gauge_set( + "lightllm_pd_master_admission_queue_size", + controller.queued_request_count, + ) + self.metric_client.gauge_set( + "lightllm_pd_master_admission_active_slots", + controller.active_slots, + ) + self.metric_client.gauge_set( + "lightllm_pd_master_admission_cold_capacity", + controller.cold_capacity, + ) + def tokens(self, prompt, multimodal_params, samping_params: SamplingParams, kwargs=None): kwargs = {} if kwargs is None else kwargs prompt_ids = self.tokenizer.encode(prompt, None, **kwargs) @@ -130,27 +171,19 @@ async def generate( multimodal_params: MultimodalParams, request: Request, ): + admission_lease = None + admission_request = None + observed_prefill_ids = set() + session_key = self._get_session_key(request) if not self.args.disable_pd_master_decode_capacity_limit: - decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.pd_manager.decode_nodes) - # 每个请求默认获得半个解码波次的缓冲容量;可复用的提示词缓存会逐步开放另外半个波次, - # 同时以两个完整解码波次作为硬上限。 - general_waiting_capacity = (decode_capacity + 1) // 2 - cache_waiting_capacity = decode_capacity // 2 - hard_admission_limit = decode_capacity + general_waiting_capacity + cache_waiting_capacity - if self.running_request_count >= hard_admission_limit: - raise ServerBusyError() - - general_admission_limit = decode_capacity + general_waiting_capacity - if self.running_request_count >= general_admission_limit: - estimated_cache_hit_rate = self.pd_manager.selector.estimate_prompt_cache_hit_rate(prompt) - if estimated_cache_hit_rate is None or not math.isfinite(estimated_cache_hit_rate): - estimated_cache_hit_rate = 0.0 - estimated_cache_hit_rate = min(max(estimated_cache_hit_rate, 0.0), 1.0) - cache_aware_admission_limit = general_admission_limit + math.ceil( - cache_waiting_capacity * estimated_cache_hit_rate - ) - if self.running_request_count >= cache_aware_admission_limit: - raise ServerBusyError() + admission_request = self._build_admission_request(prompt, sampling_params, session_key) + admission_lease = await self.admission_controller.acquire(admission_request) + self.metric_client.histogram_observe( + "lightllm_request_queue_duration_bucket", admission_lease.waited_seconds + ) + if admission_lease.waited_seconds > 0: + # 排队期间节点和缓存内容可能变化;派发前重新匹配,避免复用过期快照。 + self.pd_manager.selector.estimate_prompt_cache_hit_rate(prompt) was_idle = self.running_request_count == 0 self.running_request_count += 1 @@ -159,9 +192,62 @@ async def generate( try: async with aclosing(self._generate(prompt, sampling_params, multimodal_params, request)) as generator: async for result in generator: + if ( + admission_request is not None + and isinstance(result, tuple) + and len(result) >= 3 + and result[0] not in observed_prefill_ids + and isinstance(result[2], dict) + and "prompt_tokens" in result[2] + ): + observed_prefill_ids.add(result[0]) + self.admission_controller.record_prefill_result( + admission_request, + result[2]["prompt_tokens"], + result[2].get("prompt_cache_len", 0), + ) + if session_key is not None and not self.session_tracker.is_continuation(session_key): + self.session_tracker.mark_success(session_key) + self.admission_controller.promote_session(session_key) yield result finally: self.running_request_count -= 1 + if admission_lease is not None: + admission_lease.release() + + def _get_session_key(self, request: Optional[Request]) -> Optional[str]: + if request is None: + return None + session_key = request.headers.get("X-Session-Id", "").strip() + return session_key or None + + def _build_admission_request( + self, + prompt: Union[str, List[int]], + sampling_params: Optional[SamplingParams], + session_key: Optional[str], + ) -> AdmissionRequest: + estimated_cache_hit_rate = self.pd_manager.selector.estimate_prompt_cache_hit_rate(prompt) + if estimated_cache_hit_rate is None or not math.isfinite(estimated_cache_hit_rate): + estimated_cache_hit_rate = 0.0 + estimated_cache_hit_rate = min(max(estimated_cache_hit_rate, 0.0), 1.0) + + if self.session_tracker.is_continuation(session_key): + priority = AdmissionPriority.CONTINUATION + elif estimated_cache_hit_rate >= self.admission_policy.probable_cache_hit_threshold: + priority = AdmissionPriority.PROBABLE_CACHE_HIT + else: + priority = AdmissionPriority.COLD + + decode_slots = max(1, int(getattr(sampling_params, "n", 1) or 1)) + prompt_size = len(prompt) if prompt is not None else 0 + estimated_uncached_work = math.ceil(prompt_size * (1.0 - estimated_cache_hit_rate)) * decode_slots + return AdmissionRequest( + session_key=session_key, + priority=priority, + decode_slots=decode_slots, + estimated_uncached_work=estimated_uncached_work, + ) async def _generate( self, @@ -660,7 +746,7 @@ async def handle_loop(self): for obj in objs: if obj[0] == ObjType.TOKEN_PACKS: token_list, node_load_info = obj[1], obj[2] - self.pd_manager.update_node_load_info(node_load_info) + self.update_node_load_info(node_load_info) for sub_req_id, text, metadata, finish_status in token_list: finish_status: FinishStatus = finish_status @@ -780,6 +866,54 @@ def __init__(self, args: StartArgs): self.selector = create_selector(args.select_p_d_node_strategy, self) return + def get_decode_capacity(self) -> int: + return sum( + node.capacity_share if node.capacity_share is not None else node.start_args["running_max_req_size"] + for node in self.decode_nodes + ) + + def get_prefill_cache_capacity(self) -> Optional[CacheCapacitySnapshot]: + """汇总当前 Master 对应的 Prefill cache 份额;遥测不完整时不参与限流。""" + if not self.prefill_nodes: + return None + + statuses = [node.run_status for node in self.prefill_nodes] + if any(status.radix_cache_capacity_tokens <= 0 or status.report_time <= 0 for status in statuses): + return None + + full_decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.decode_nodes) + local_decode_capacity = self.get_decode_capacity() + if full_decode_capacity <= 0 or local_decode_capacity <= 0: + return None + + # 所有 Master 都能看到同一组 P 节点,因此按本 Master 的 Decode 租约比例 + # 切分缓存余量,避免每个 Master 重复消费整份 headroom。 + share_ratio = min(1.0, local_decode_capacity / full_decode_capacity) + # total_token_usage_rate 已排除可驱逐 Radix token,因此下面两项分别表示 + # 当前可驱逐缓存量和运行中请求之外还能留给缓存的容量。 + total_tokens = int( + sum( + max(0, status.radix_cache_total_tokens - status.radix_cache_refed_tokens) + for status in statuses + ) + * share_ratio + ) + capacity_tokens = int( + sum( + max( + 0, + status.radix_cache_capacity_tokens + * (1.0 - min(max(status.total_token_usage_rate, 0.0), 1.0)), + ) + for status in statuses + ) + * share_ratio + ) + return CacheCapacitySnapshot( + total_tokens=total_tokens, + capacity_tokens=capacity_tokens, + ) + def is_pd_nodes_ready(self): prefill_node_count = len(self.prefill_nodes) decode_node_count = len(self.decode_nodes) @@ -894,10 +1028,25 @@ def update_node_load_info(self, load_info: Optional[dict]): if load_info is None: return client_ip_port = load_info["client_ip_port"] - total_token_usage_rate = load_info["total_token_usage_rate"] pd_client = self.url_to_pd_nodes.get(client_ip_port) - pd_client.run_status.total_token_usage_rate = total_token_usage_rate - except BaseException as e: + if pd_client is None: + return + pd_client.run_status.total_token_usage_rate = load_info["total_token_usage_rate"] + pd_client.run_status.radix_cache_total_tokens = load_info.get("radix_cache_total_tokens", 0) + pd_client.run_status.radix_cache_refed_tokens = load_info.get("radix_cache_refed_tokens", 0) + pd_client.run_status.radix_cache_capacity_tokens = load_info.get("radix_cache_capacity_tokens", 0) + pd_client.run_status.report_time = time.monotonic() + + capacity_epoch = int(load_info.get("capacity_epoch", pd_client.capacity_epoch)) + if capacity_epoch >= pd_client.capacity_epoch: + fallback_capacity = ( + pd_client.capacity_share + if pd_client.capacity_share is not None + else pd_client.start_args["running_max_req_size"] + ) + pd_client.capacity_share = max(0, int(load_info.get("capacity_share", fallback_capacity))) + pd_client.capacity_epoch = capacity_epoch + except Exception as e: logger.warning(f"udpate node load info failed, load_info: {load_info} error: {str(e)}") return diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index 0d42462c3..d11794c6a 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -32,6 +32,9 @@ "lightllm_cache_hit_rate": "Prefix cache hit rate of latest completed request", "lightllm_gen_throughput": "Generation throughput of latest completed request (tokens/s)", "lightllm_num_running_reqs": "Number of running requests", + "lightllm_pd_master_admission_queue_size": "Number of requests waiting at the PD master admission queue", + "lightllm_pd_master_admission_active_slots": "Number of decode slots leased by the PD master", + "lightllm_pd_master_admission_cold_capacity": "Current cold-request slot capacity at the PD master", } @@ -111,6 +114,9 @@ def init_metrics(self, args): self.create_gauge("lightllm_cache_hit_rate") self.create_gauge("lightllm_gen_throughput") self.create_gauge("lightllm_num_running_reqs") + self.create_gauge("lightllm_pd_master_admission_queue_size") + self.create_gauge("lightllm_pd_master_admission_active_slots") + self.create_gauge("lightllm_pd_master_admission_cold_capacity") def create_histogram(self, name, buckets, labelnames=None): all_labels = ["model_name"] + (labelnames or []) diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index 8a1e2bd42..e36dc1e67 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -47,6 +47,10 @@ class ObjType(enum.Enum): @dataclass class _PD_Client_RunStatus: total_token_usage_rate: float = 0.0 # pd 节点上的 token 使用率 + radix_cache_total_tokens: int = 0 + radix_cache_refed_tokens: int = 0 + radix_cache_capacity_tokens: int = 0 + report_time: float = 0.0 @dataclass @@ -57,6 +61,9 @@ class PD_Client_Obj: start_args: object # 节点的启动参数信息,用于做匹配性的校验,防止运行过程中出现问题。 websocket: WebSocket = None # 用于通信的 websocket 连接对象 run_status: _PD_Client_RunStatus = field(default_factory=_PD_Client_RunStatus) + # 节点租给当前 PD Master 的请求槽位;多 Master 之间的份额互不重叠。 + capacity_share: Optional[int] = None + capacity_epoch: int = 0 # cache-aware 选点用:当前派发到该节点且尚未产出首 token 的 prompt 字符数。 dispatched_prompt_chars: int = 0 # 当前派发到该节点且尚未产出首 token 的请求数。 @@ -67,6 +74,10 @@ def __post_init__(self): error_info = f"""mode must in ["prefill", "decode"], but get {self.mode}""" logger.error(error_info) raise ValueError(error_info) + if self.capacity_share is not None and self.capacity_share < 0: + raise ValueError("capacity_share must be non-negative") + if self.capacity_epoch < 0: + raise ValueError("capacity_epoch must be non-negative") return def to_llm_url(self): diff --git a/unit_tests/server/test_pd_admission.py b/unit_tests/server/test_pd_admission.py new file mode 100644 index 000000000..92d0760a0 --- /dev/null +++ b/unit_tests/server/test_pd_admission.py @@ -0,0 +1,333 @@ +import asyncio + +import pytest + +from lightllm.server.httpserver_for_pd_master.admission import ( + AdmissionPolicy, + AdmissionPriority, + AdmissionRequest, + CacheCapacitySnapshot, + PDAdmissionController, + SessionTracker, +) +from lightllm.utils.error_utils import ServerBusyError + + +def _request( + priority=AdmissionPriority.COLD, + session_key=None, + decode_slots=1, + estimated_uncached_work=0, +): + return AdmissionRequest( + session_key=session_key, + priority=priority, + decode_slots=decode_slots, + estimated_uncached_work=estimated_uncached_work, + ) + + +def test_admission_waits_before_granting_more_than_decode_capacity(): + async def run(): + controller = PDAdmissionController(lambda: 1) + first = await controller.acquire(_request()) + second_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + + assert controller.active_slots == 1 + assert controller.queued_slots == 1 + assert second_task.done() is False + + first.release() + second = await second_task + assert controller.active_slots == 1 + assert controller.queued_slots == 0 + assert second.waited_seconds >= 0 + second.release() + + asyncio.run(run()) + + +def test_admission_prioritizes_continuations_without_starving_lower_classes(): + async def run(): + controller = PDAdmissionController(lambda: 3) + active = [await controller.acquire(_request()) for _ in range(3)] + cold_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.COLD))) + probable_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.PROBABLE_CACHE_HIT))) + continuation_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + + active[0].release() + continuation = await continuation_task + assert probable_task.done() is False + assert cold_task.done() is False + + active[1].release() + probable = await probable_task + active[2].release() + cold = await cold_task + + continuation.release() + probable.release() + cold.release() + + asyncio.run(run()) + + +def test_higher_priority_request_can_replace_a_queued_cold_request(): + async def run(): + controller = PDAdmissionController(lambda: 1) + active = await controller.acquire(_request()) + cold_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.COLD))) + await asyncio.sleep(0) + continuation_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + + with pytest.raises(ServerBusyError, match="Superseded"): + await cold_task + + active.release() + continuation = await continuation_task + continuation.release() + + asyncio.run(run()) + + +def test_multi_choice_request_acquires_all_slots_atomically(): + async def run(): + controller = PDAdmissionController(lambda: 3) + active = await controller.acquire(_request(decode_slots=2)) + multi_choice_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) + await asyncio.sleep(0) + + assert controller.active_slots == 2 + assert multi_choice_task.done() is False + + active.release() + multi_choice = await multi_choice_task + assert controller.active_slots == 2 + multi_choice.release() + + asyncio.run(run()) + + +def test_multi_choice_request_reserves_capacity_across_individual_releases(): + async def run(): + controller = PDAdmissionController(lambda: 3) + active = [await controller.acquire(_request()) for _ in range(3)] + multi_choice_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) + later_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + + active[0].release() + await asyncio.sleep(0) + assert multi_choice_task.done() is False + assert later_task.done() is False + + active[1].release() + multi_choice = await multi_choice_task + assert later_task.done() is False + + active[2].release() + later = await later_task + multi_choice.release() + later.release() + + asyncio.run(run()) + + +def test_same_session_is_fifo_while_other_sessions_can_make_progress(): + async def run(): + controller = PDAdmissionController(lambda: 2) + first_a = await controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="a")) + second_a_task = asyncio.create_task( + controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="a")) + ) + await asyncio.sleep(0) + session_b = await controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="b")) + + assert second_a_task.done() is False + session_b.release() + assert second_a_task.done() is False + + first_a.release() + second_a = await second_a_task + second_a.release() + + asyncio.run(run()) + + +def test_cancelled_waiter_is_removed_and_does_not_leak_capacity(): + async def run(): + controller = PDAdmissionController(lambda: 1) + active = await controller.acquire(_request()) + waiting_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + + waiting_task.cancel() + with pytest.raises(asyncio.CancelledError): + await waiting_task + + assert controller.queued_request_count == 0 + assert controller.queued_slots == 0 + active.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_wait_timeout_removes_request_from_queue(): + async def run(): + policy = AdmissionPolicy( + continuation_max_wait_seconds=0.01, + probable_cache_hit_max_wait_seconds=0.01, + cold_max_wait_seconds=0.01, + ) + controller = PDAdmissionController(lambda: 1, policy=policy) + active = await controller.acquire(_request()) + + with pytest.raises(ServerBusyError, match="wait timed out"): + await controller.acquire(_request()) + + assert controller.queued_request_count == 0 + active.release() + + asyncio.run(run()) + + +def test_capacity_changes_wake_waiters_without_overcommitting(): + async def run(): + capacity = [1] + controller = PDAdmissionController(lambda: capacity[0]) + first = await controller.acquire(_request()) + second_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + + capacity[0] = 2 + controller.on_capacity_change() + second = await second_task + assert controller.active_slots == 2 + + capacity[0] = 1 + controller.on_capacity_change() + third_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + assert third_task.done() is False + + first.release() + await asyncio.sleep(0) + assert third_task.done() is False + second.release() + third = await third_task + third.release() + + asyncio.run(run()) + + +def test_session_tracker_requires_a_recent_observed_success(): + now = [0.0] + tracker = SessionTracker(ttl_seconds=10, max_sessions=2, clock=lambda: now[0]) + + assert tracker.is_continuation("session-a") is False + tracker.mark_success("session-a") + assert tracker.is_continuation("session-a") is True + + now[0] = 11 + assert tracker.is_continuation("session-a") is False + + tracker.mark_success("session-a") + tracker.mark_success("session-b") + tracker.mark_success("session-c") + assert tracker.is_continuation("session-a") is False + assert tracker.is_continuation("session-b") is True + assert tracker.is_continuation("session-c") is True + + +def test_successful_session_promotes_its_waiting_requests(): + async def run(): + controller = PDAdmissionController( + lambda: 1, + policy=AdmissionPolicy(waiting_decode_waves=2), + ) + active = await controller.acquire(_request()) + same_session_task = asyncio.create_task(controller.acquire(_request(session_key="session-a"))) + other_cold_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + + controller.promote_session("session-a") + active.release() + promoted = await same_session_task + assert promoted.request.priority == AdmissionPriority.CONTINUATION + assert other_cold_task.done() is False + + promoted.release() + other = await other_cold_task + other.release() + + asyncio.run(run()) + + +def test_cold_capacity_tracks_live_cache_headroom_and_actual_miss_size(): + snapshot = [CacheCapacitySnapshot(total_tokens=200, capacity_tokens=1000)] + controller = PDAdmissionController( + lambda: 4, + cache_capacity_provider=lambda: snapshot[0], + ) + cold = _request(AdmissionPriority.COLD) + + controller.record_prefill_result(cold, prompt_tokens=100, cached_tokens=0) + assert controller.cold_capacity == 4 + + snapshot[0] = CacheCapacitySnapshot(total_tokens=800, capacity_tokens=1000) + controller.on_capacity_change() + assert controller.cold_capacity == 2 + + snapshot[0] = CacheCapacitySnapshot(total_tokens=1000, capacity_tokens=1000) + controller.on_capacity_change() + assert controller.cold_capacity == 1 + + +def test_cache_friendly_request_bypasses_cold_capacity(): + async def run(): + snapshot = CacheCapacitySnapshot(total_tokens=1000, capacity_tokens=1000) + controller = PDAdmissionController( + lambda: 3, + cache_capacity_provider=lambda: snapshot, + ) + cold_request = _request(AdmissionPriority.COLD) + controller.record_prefill_result(cold_request, prompt_tokens=100, cached_tokens=0) + + first_cold = await controller.acquire(cold_request) + second_cold_task = asyncio.create_task(controller.acquire(cold_request)) + await asyncio.sleep(0) + assert second_cold_task.done() is False + + probable = await controller.acquire(_request(AdmissionPriority.PROBABLE_CACHE_HIT)) + assert controller.active_slots == 2 + probable.release() + + first_cold.release() + second_cold = await second_cold_task + second_cold.release() + + asyncio.run(run()) + + +def test_smaller_cold_request_is_dispatched_first(): + async def run(): + controller = PDAdmissionController(lambda: 2) + active = [await controller.acquire(_request()) for _ in range(2)] + large_task = asyncio.create_task(controller.acquire(_request(estimated_uncached_work=100))) + small_task = asyncio.create_task(controller.acquire(_request(estimated_uncached_work=10))) + await asyncio.sleep(0) + + active[0].release() + small = await small_task + assert large_task.done() is False + + active[1].release() + large = await large_task + small.release() + large.release() + + asyncio.run(run()) diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index b7e06b140..f3467fa70 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -6,8 +6,14 @@ from easydict import EasyDict from lightllm.server.core.objs.start_args_type import StartArgs +from lightllm.server.httpserver.pd_loop import _allocate_capacity_share, _update_pd_master_membership +from lightllm.server.httpserver_for_pd_master.admission import ( + AdmissionPolicy, + AdmissionPriority, + PDAdmissionController, + SessionTracker, +) from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager -from lightllm.utils.error_utils import ServerBusyError def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): @@ -185,6 +191,134 @@ def test_pd_manager_without_connected_nodes_is_healthy(): assert asyncio.run(manager.check_pd_nodes_health()) is True +def test_pd_node_capacity_is_partitioned_without_overlap(): + master_ids = [30, 10, 20] + shares = [_allocate_capacity_share(8, master_ids, node_id) for node_id in sorted(master_ids)] + + assert shares == [3, 3, 2] + assert sum(shares) == 8 + assert _allocate_capacity_share(8, master_ids, 99) == 0 + + +def test_pd_master_membership_change_advances_epoch_and_wakes_heartbeats(monkeypatch): + async def run(): + manager = SimpleNamespace() + timestamps = iter([100, 200]) + monkeypatch.setattr("lightllm.server.httpserver.pd_loop.time.time_ns", lambda: next(timestamps)) + + _update_pd_master_membership(manager, {20: object(), 10: object()}) + assert manager.pd_master_ids == (10, 20) + assert manager.pd_master_capacity_epoch == 100 + assert manager.pd_master_membership_changed.is_set() + + manager.pd_master_membership_changed.clear() + _update_pd_master_membership(manager, {10: object(), 20: object()}) + assert manager.pd_master_membership_changed.is_set() is False + + _update_pd_master_membership(manager, {10: object()}) + assert manager.pd_master_capacity_epoch == 200 + assert manager.pd_master_membership_changed.is_set() + + asyncio.run(run()) + + +def test_pd_manager_uses_latest_decode_capacity_lease_and_cache_telemetry(): + args = StartArgs() + manager = PDManager(args) + client_ip_port = "10.0.0.2:8000" + manager.register_pd( + { + "node_id": 2, + "client_ip_port": client_ip_port, + "mode": "decode", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + }, + "capacity_share": 3, + "capacity_epoch": 100, + }, + websocket=object(), + ) + + assert manager.get_decode_capacity() == 3 + + manager.update_node_load_info( + { + "client_ip_port": client_ip_port, + "total_token_usage_rate": 0.25, + "capacity_share": 1, + "capacity_epoch": 99, + } + ) + assert manager.get_decode_capacity() == 3 + + manager.update_node_load_info( + { + "client_ip_port": client_ip_port, + "total_token_usage_rate": 0.5, + "capacity_share": 2, + "capacity_epoch": 101, + "radix_cache_total_tokens": 700, + "radix_cache_refed_tokens": 200, + "radix_cache_capacity_tokens": 1000, + } + ) + node = manager.decode_nodes[0] + assert manager.get_decode_capacity() == 2 + assert node.run_status.total_token_usage_rate == 0.5 + assert node.run_status.radix_cache_total_tokens == 700 + assert node.run_status.radix_cache_refed_tokens == 200 + assert node.run_status.radix_cache_capacity_tokens == 1000 + + +def test_prefill_cache_headroom_is_scaled_to_the_master_decode_lease(): + args = StartArgs() + manager = PDManager(args) + manager.register_pd( + { + "node_id": 1, + "client_ip_port": "10.0.0.1:8000", + "mode": "prefill", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + "max_image_pixels": args.max_image_pixels, + "disable_image_resize": args.disable_image_resize, + }, + }, + websocket=object(), + ) + manager.register_pd( + { + "node_id": 2, + "client_ip_port": "10.0.0.2:8000", + "mode": "decode", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + }, + "capacity_share": 4, + }, + websocket=object(), + ) + manager.update_node_load_info( + { + "client_ip_port": "10.0.0.1:8000", + "total_token_usage_rate": 0.25, + "radix_cache_total_tokens": 800, + "radix_cache_refed_tokens": 100, + "radix_cache_capacity_tokens": 1000, + } + ) + + snapshot = manager.get_prefill_cache_capacity() + assert snapshot is not None + assert snapshot.total_tokens == 350 + assert snapshot.capacity_tokens == 375 + assert snapshot.free_tokens == 25 + + def test_prefill_registration_preserves_existing_inflight_prompt_chars(): args = StartArgs() manager = PDManager(args) @@ -263,61 +397,72 @@ def test_pd_master_inference_health_matches_normal_node_semantics(monkeypatch): assert HttpServerManagerForPDMaster.is_healthy(manager) is True -@pytest.mark.parametrize( - ("decode_capacity", "estimated_cache_hit_rate", "running_request_count", "is_rejected"), - [ - (8, None, 11, False), - (8, None, 12, True), - (3, None, 4, False), - (3, None, 5, True), - (8, 0.0, 12, True), - (8, 0.5, 13, False), - (8, 0.5, 14, True), - (8, 1.0, 15, False), - (8, 1.0, 16, True), - ], -) -def test_pd_master_admission_adapts_to_capacity_and_cache_hit_rate( - decode_capacity, estimated_cache_hit_rate, running_request_count, is_rejected -): - manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) - manager.args = StartArgs() - estimate_calls = [] +def test_pd_master_waits_before_dispatching_beyond_decode_capacity(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs() + manager.pd_manager = SimpleNamespace( + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: 0.0), + ) + manager.admission_policy = AdmissionPolicy() + manager.admission_controller = PDAdmissionController(lambda: 1, policy=manager.admission_policy) + manager.session_tracker = SessionTracker( + ttl_seconds=manager.admission_policy.active_session_ttl_seconds, + max_sessions=manager.admission_policy.max_tracked_sessions, + ) + manager.metric_client = SimpleNamespace(histogram_observe=lambda *_args: None) + manager.running_request_count = 0 + manager.latest_success_infer_time = 0 + dispatched_prompts = [] - def estimate_prompt_cache_hit_rate(prompt): - estimate_calls.append(prompt) - return estimated_cache_hit_rate + async def fake_generate(prompt, *_args): + dispatched_prompts.append(prompt) + yield prompt - manager.pd_manager = SimpleNamespace( - decode_nodes=[SimpleNamespace(start_args={"running_max_req_size": decode_capacity})], - selector=SimpleNamespace(estimate_prompt_cache_hit_rate=estimate_prompt_cache_hit_rate), - ) - manager.running_request_count = running_request_count - manager.latest_success_infer_time = 0 + manager._generate = fake_generate + first = manager.generate("first", None, None, None) + assert await first.__anext__() == "first" - async def fake_generate(*_args): - yield "result" + second = manager.generate("second", None, None, None) + second_result = asyncio.create_task(second.__anext__()) + await asyncio.sleep(0) + assert dispatched_prompts == ["first"] + assert manager.admission_controller.queued_request_count == 1 - manager._generate = fake_generate + await first.aclose() + assert await second_result == "second" + assert dispatched_prompts == ["first", "second"] + await second.aclose() + assert manager.admission_controller.active_slots == 0 - async def consume_one_result(): - generator = manager.generate("multi-turn prompt", None, None, None) - try: - assert await generator.__anext__() == "result" - finally: - await generator.aclose() - - if is_rejected: - with pytest.raises(ServerBusyError): - asyncio.run(consume_one_result()) - assert manager.running_request_count == running_request_count - else: - asyncio.run(consume_one_result()) - assert manager.running_request_count == running_request_count - - general_admission_limit = decode_capacity + (decode_capacity + 1) // 2 - should_estimate_cache = general_admission_limit <= running_request_count < 2 * decode_capacity - assert len(estimate_calls) == int(should_estimate_cache) + asyncio.run(run()) + + +def test_pd_master_admission_classifies_session_cache_and_multi_choice_cost(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.pd_manager = SimpleNamespace( + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: 0.75), + ) + manager.admission_policy = AdmissionPolicy() + manager.session_tracker = SessionTracker(ttl_seconds=60, max_sessions=10) + sampling_params = SimpleNamespace(n=3) + + probable = manager._build_admission_request( + "abcdefghij", + sampling_params, + session_key="session-a", + ) + assert probable.priority == AdmissionPriority.PROBABLE_CACHE_HIT + assert probable.decode_slots == 3 + assert probable.estimated_uncached_work == 9 + + manager.session_tracker.mark_success("session-a") + continuation = manager._build_admission_request( + "abcdefghij", + sampling_params, + session_key="session-a", + ) + assert continuation.priority == AdmissionPriority.CONTINUATION def test_pd_master_restores_request_count_when_preload_fails(): From c16db61baa8b916f035d9575e71ad8348df7e62c Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 31 Aug 2026 22:40:03 +0800 Subject: [PATCH 6/9] style(pd): apply pre-commit formatting --- lightllm/server/httpserver_for_pd_master/manager.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index ea26bd328..0624b44be 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -892,18 +892,14 @@ def get_prefill_cache_capacity(self) -> Optional[CacheCapacitySnapshot]: # total_token_usage_rate 已排除可驱逐 Radix token,因此下面两项分别表示 # 当前可驱逐缓存量和运行中请求之外还能留给缓存的容量。 total_tokens = int( - sum( - max(0, status.radix_cache_total_tokens - status.radix_cache_refed_tokens) - for status in statuses - ) + sum(max(0, status.radix_cache_total_tokens - status.radix_cache_refed_tokens) for status in statuses) * share_ratio ) capacity_tokens = int( sum( max( 0, - status.radix_cache_capacity_tokens - * (1.0 - min(max(status.total_token_usage_rate, 0.0), 1.0)), + status.radix_cache_capacity_tokens * (1.0 - min(max(status.total_token_usage_rate, 0.0), 1.0)), ) for status in statuses ) From 81486f67e75659087a5d7e150001f497947decf2 Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 31 Aug 2026 23:27:19 +0800 Subject: [PATCH 7/9] fix(pd): close admission priority feedback gaps --- .../httpserver_for_pd_master/admission.py | 125 +++++++++++++----- .../httpserver_for_pd_master/manager.py | 1 + unit_tests/server/test_pd_admission.py | 94 +++++++++++++ 3 files changed, 190 insertions(+), 30 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/admission.py b/lightllm/server/httpserver_for_pd_master/admission.py index 6cda68274..94ea62e53 100644 --- a/lightllm/server/httpserver_for_pd_master/admission.py +++ b/lightllm/server/httpserver_for_pd_master/admission.py @@ -151,6 +151,8 @@ class _WaitingRequest: sequence_id: int request: AdmissionRequest enqueue_time: float + deadline: float + deadline_changed: asyncio.Event future: asyncio.Future @@ -162,10 +164,12 @@ def __init__( controller: "PDAdmissionController", request: AdmissionRequest, waited_seconds: float, + cold_slots: int, ) -> None: self._controller = controller self.request = request self.waited_seconds = waited_seconds + self._cold_slots = cold_slots self._released = False async def __aenter__(self) -> "AdmissionLease": @@ -206,6 +210,7 @@ def __init__( self._active_slots = 0 self._active_cold_slots = 0 self._average_cold_uncached_tokens: Optional[float] = None + self._probable_actual_hit_rate: Optional[float] = None self._active_sessions = set() self._queues: Dict[AdmissionPriority, Deque[_WaitingRequest]] = { priority: deque() for priority in self._PRIORITY_ORDER @@ -277,10 +282,13 @@ async def acquire(self, request: AdmissionRequest) -> AdmissionLease: return lease loop = asyncio.get_running_loop() + enqueue_time = self._clock() waiter = _WaitingRequest( sequence_id=self._sequence_id, request=request, - enqueue_time=self._clock(), + enqueue_time=enqueue_time, + deadline=enqueue_time + self.policy.max_wait_seconds(request.priority), + deadline_changed=asyncio.Event(), future=loop.create_future(), ) self._sequence_id += 1 @@ -292,10 +300,7 @@ async def acquire(self, request: AdmissionRequest) -> AdmissionLease: self._drain() try: - return await asyncio.wait_for( - asyncio.shield(waiter.future), - timeout=self.policy.max_wait_seconds(request.priority), - ) + return await self._wait_for_lease(waiter) except asyncio.TimeoutError as exc: lease = self._cancel_waiter_or_take_lease(waiter) if lease is not None: @@ -316,21 +321,27 @@ def record_prefill_result( prompt_tokens: int, cached_tokens: int, ) -> None: - """用冷请求的真实未命中量更新下一轮冷容量。""" - if request.priority != AdmissionPriority.COLD: + """用真实命中结果更新预计命中的可信度和冷请求容量。""" + if request.priority == AdmissionPriority.CONTINUATION: return prompt_tokens = max(0, int(prompt_tokens)) cached_tokens = min(prompt_tokens, max(0, int(cached_tokens))) uncached_tokens = prompt_tokens - cached_tokens - # 一个 Decode 波次作为自适应窗口:容量越大,单个样本对均值的影响越小。 - sample_window = max(1, self._capacity()) - alpha = 2.0 / (sample_window + 1.0) - if self._average_cold_uncached_tokens is None: - self._average_cold_uncached_tokens = float(uncached_tokens) - else: - self._average_cold_uncached_tokens += alpha * (uncached_tokens - self._average_cold_uncached_tokens) + if request.priority == AdmissionPriority.PROBABLE_CACHE_HIT: + actual_hit_rate = cached_tokens / max(prompt_tokens, 1) + self._probable_actual_hit_rate = self._update_average( + self._probable_actual_hit_rate, + actual_hit_rate, + ) + + # 预计命中的真实命中率低于承诺阈值时,让后续同类请求也消费冷槽位。 + if self._requires_cold_capacity(request): + self._average_cold_uncached_tokens = self._update_average( + self._average_cold_uncached_tokens, + float(uncached_tokens), + ) self._drain() def promote_session(self, session_key: Optional[str]) -> None: @@ -347,11 +358,33 @@ def promote_session(self, session_key: Optional[str]) -> None: continue self._queues[old_priority].remove(waiter) waiter.request = replace(waiter.request, priority=AdmissionPriority.CONTINUATION) + waiter.deadline = max( + waiter.deadline, + waiter.enqueue_time + self.policy.continuation_max_wait_seconds, + ) + waiter.deadline_changed.set() self._queues[AdmissionPriority.CONTINUATION].append(waiter) self._drain() + def _update_average(self, current: Optional[float], sample: float) -> float: + # 一个 Decode 波次作为自适应窗口:容量越大,单个样本对均值的影响越小。 + sample_window = max(1, self._capacity()) + alpha = 2.0 / (sample_window + 1.0) + if current is None: + return sample + return current + alpha * (sample - current) + + def _requires_cold_capacity(self, request: AdmissionRequest) -> bool: + if request.priority == AdmissionPriority.COLD: + return True + return ( + request.priority == AdmissionPriority.PROBABLE_CACHE_HIT + and self._probable_actual_hit_rate is not None + and self._probable_actual_hit_rate < self.policy.probable_cache_hit_threshold + ) + def _has_cold_capacity(self, request: AdmissionRequest) -> bool: - if request.priority != AdmissionPriority.COLD: + if not self._requires_cold_capacity(request): return True return self._active_cold_slots + request.decode_slots <= self.cold_capacity @@ -364,17 +397,16 @@ def _can_activate(self, request: AdmissionRequest) -> bool: def _activate(self, request: AdmissionRequest, waited_seconds: float) -> AdmissionLease: self._active_slots += request.decode_slots - if request.priority == AdmissionPriority.COLD: - self._active_cold_slots += request.decode_slots + cold_slots = request.decode_slots if self._requires_cold_capacity(request) else 0 + self._active_cold_slots += cold_slots if request.session_key is not None: self._active_sessions.add(request.session_key) - return AdmissionLease(self, request, waited_seconds) + return AdmissionLease(self, request, waited_seconds, cold_slots) def _release(self, lease: AdmissionLease) -> None: request = lease.request self._active_slots -= request.decode_slots - if request.priority == AdmissionPriority.COLD: - self._active_cold_slots -= request.decode_slots + self._active_cold_slots -= lease._cold_slots if self._active_slots < 0: raise RuntimeError("PD admission active slot count became negative") if self._active_cold_slots < 0: @@ -452,14 +484,39 @@ def _cancel_waiter_or_take_lease(self, waiter: _WaitingRequest) -> Optional[Admi return result return None + async def _wait_for_lease(self, waiter: _WaitingRequest) -> AdmissionLease: + while True: + deadline_changed_task = asyncio.create_task(waiter.deadline_changed.wait()) + try: + done, _ = await asyncio.wait( + (waiter.future, deadline_changed_task), + timeout=max(0.0, waiter.deadline - self._clock()), + return_when=asyncio.FIRST_COMPLETED, + ) + finally: + if not deadline_changed_task.done(): + deadline_changed_task.cancel() + + if waiter.future in done: + return waiter.future.result() + if deadline_changed_task in done: + waiter.deadline_changed.clear() + continue + raise asyncio.TimeoutError + + def _session_is_grantable(self, waiter: _WaitingRequest) -> bool: + session_key = waiter.request.session_key + if session_key is None: + return True + if session_key in self._active_sessions: + return False + return self._session_queues[session_key][0] is waiter + def _first_grantable(self, priority: AdmissionPriority) -> Optional[_WaitingRequest]: candidates = [] available_decode_slots = self._capacity() - self._active_slots for waiter in self._queues[priority]: - session_key = waiter.request.session_key - if session_key is not None and session_key in self._active_sessions: - continue - if session_key is not None and self._session_queues[session_key][0] is not waiter: + if not self._session_is_grantable(waiter): continue if priority != AdmissionPriority.COLD: return waiter @@ -480,6 +537,15 @@ def _first_grantable(self, priority: AdmissionPriority) -> Optional[_WaitingRequ ), ) + def _first_fitting_higher_priority(self, blocked_priority: AdmissionPriority) -> Optional[_WaitingRequest]: + for priority in self._PRIORITY_ORDER: + if priority <= blocked_priority: + continue + for waiter in self._queues[priority]: + if self._session_is_grantable(waiter) and self._can_activate(waiter.request): + return waiter + return None + def _select_next(self) -> Optional[_WaitingRequest]: schedule_size = len(self._schedule) for offset in range(schedule_size): @@ -496,13 +562,12 @@ def _drain(self) -> None: waiter = self._select_next() if waiter is None: break - if self._active_slots + waiter.request.decode_slots > self._capacity(): - # 为需要多个 choice slot 的老请求保留逐步释放出来的容量,避免永久饥饿。 + if not self._can_activate(waiter.request): + # 为被选中的多 choice 请求积累槽位,但不因此阻塞当前可以执行的更高优先级请求。 self._schedule_index = schedule_index - break - if not self._has_cold_capacity(waiter.request): - self._schedule_index = schedule_index - break + waiter = self._first_fitting_higher_priority(waiter.request.priority) + if waiter is None: + break if not self._remove_waiter(waiter): continue lease = self._activate( diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 0624b44be..6bf32108b 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -178,6 +178,7 @@ async def generate( if not self.args.disable_pd_master_decode_capacity_limit: admission_request = self._build_admission_request(prompt, sampling_params, session_key) admission_lease = await self.admission_controller.acquire(admission_request) + admission_request = admission_lease.request self.metric_client.histogram_observe( "lightllm_request_queue_duration_bucket", admission_lease.waited_seconds ) diff --git a/unit_tests/server/test_pd_admission.py b/unit_tests/server/test_pd_admission.py index 92d0760a0..80a623955 100644 --- a/unit_tests/server/test_pd_admission.py +++ b/unit_tests/server/test_pd_admission.py @@ -136,6 +136,40 @@ async def run(): asyncio.run(run()) +def test_multi_choice_reservation_does_not_block_fitting_higher_priority_request(): + async def run(): + controller = PDAdmissionController(lambda: 3) + active = [await controller.acquire(_request()) for _ in range(3)] + + # 先消费调度表中的 continuation 和 probable 配额,使下一次轮到 cold。 + first_continuation_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + first_probable_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.PROBABLE_CACHE_HIT))) + await asyncio.sleep(0) + active[0].release() + first_continuation = await first_continuation_task + active[1].release() + first_probable = await first_probable_task + + large_cold_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) + later_continuation_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + + active[2].release() + later_continuation = await later_continuation_task + assert controller.active_slots == 3 + assert large_cold_task.done() is False + + first_continuation.release() + assert large_cold_task.done() is False + first_probable.release() + large_cold = await large_cold_task + + later_continuation.release() + large_cold.release() + + asyncio.run(run()) + + def test_same_session_is_fifo_while_other_sessions_can_make_progress(): async def run(): controller = PDAdmissionController(lambda: 2) @@ -267,6 +301,30 @@ async def run(): asyncio.run(run()) +def test_session_promotion_extends_the_wait_deadline(): + async def run(): + policy = AdmissionPolicy( + continuation_max_wait_seconds=0.2, + probable_cache_hit_max_wait_seconds=0.1, + cold_max_wait_seconds=0.02, + ) + controller = PDAdmissionController(lambda: 1, policy=policy) + active = await controller.acquire(_request()) + waiting_task = asyncio.create_task(controller.acquire(_request(session_key="session-a"))) + await asyncio.sleep(0.005) + + controller.promote_session("session-a") + await asyncio.sleep(0.025) + assert waiting_task.done() is False + + active.release() + promoted = await waiting_task + assert promoted.request.priority == AdmissionPriority.CONTINUATION + promoted.release() + + asyncio.run(run()) + + def test_cold_capacity_tracks_live_cache_headroom_and_actual_miss_size(): snapshot = [CacheCapacitySnapshot(total_tokens=200, capacity_tokens=1000)] controller = PDAdmissionController( @@ -313,6 +371,42 @@ async def run(): asyncio.run(run()) +def test_probable_cache_hits_consume_cold_capacity_when_actual_hits_are_low(): + async def run(): + snapshot = CacheCapacitySnapshot(total_tokens=1000, capacity_tokens=1000) + controller = PDAdmissionController( + lambda: 3, + cache_capacity_provider=lambda: snapshot, + ) + probable_request = _request(AdmissionPriority.PROBABLE_CACHE_HIT) + + initially_trusted = await controller.acquire(probable_request) + assert controller.active_cold_slots == 0 + controller.record_prefill_result( + probable_request, + prompt_tokens=100, + cached_tokens=0, + ) + assert controller.cold_capacity == 1 + + first_gated = await controller.acquire(probable_request) + assert controller.active_cold_slots == 1 + second_gated_task = asyncio.create_task(controller.acquire(probable_request)) + await asyncio.sleep(0) + assert second_gated_task.done() is False + + # 可信度变化不能让已经取得的租约在释放时误扣冷槽位。 + initially_trusted.release() + assert controller.active_cold_slots == 1 + assert second_gated_task.done() is False + + first_gated.release() + second_gated = await second_gated_task + second_gated.release() + + asyncio.run(run()) + + def test_smaller_cold_request_is_dispatched_first(): async def run(): controller = PDAdmissionController(lambda: 2) From c77ab740bf690f826d62d501596c49bb0cc14721 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 1 Sep 2026 01:19:36 +0800 Subject: [PATCH 8/9] docs(pd): add Chinese comments for admission helpers --- lightllm/server/httpserver/pd_loop.py | 5 +++ .../httpserver_for_pd_master/admission.py | 40 +++++++++++++++++++ .../httpserver_for_pd_master/manager.py | 5 +++ .../pd_selector/cache_aware.py | 1 + .../pd_selector/pd_selector.py | 1 + 5 files changed, 52 insertions(+) diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index 61656459a..c3e3110b4 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -32,6 +32,7 @@ def _update_pd_master_membership(manager: HttpServerManager, pd_master_ids) -> None: + """更新 Master 成员和容量版本,并立即唤醒心跳。""" pd_master_ids = tuple(sorted(pd_master_ids)) if getattr(manager, "pd_master_ids", ()) == pd_master_ids: return @@ -56,6 +57,7 @@ def _allocate_capacity_share(total_capacity: int, pd_master_ids, pd_master_node_ def _get_radix_cache_info(): + """读取本节点各 DP 的 Radix cache token 统计。""" global _radix_cache_client, _radix_cache_client_key from lightllm.server.api_http import g_objs @@ -344,6 +346,7 @@ async def _up_tokens_to_pd_master( websocket: ClientConnection, pd_master_node_id: int, ): + """批量向 PD Master 转发生成结果和最新负载。""" while True: handle_list = await forwarding_queue.wait_to_get_all_data() @@ -357,6 +360,7 @@ async def _send_heartbeat_to_pd_master( websocket: ClientConnection, pd_master_node_id: int, ): + """定时或在成员变化时向 PD Master 上报心跳。""" heartbeat_interval_seconds = 15 membership_changed = manager.pd_master_membership_changed while True: @@ -371,6 +375,7 @@ async def _send_heartbeat_to_pd_master( # 获取节点负载信息 def _get_load_info(pd_master_node_id: int) -> dict: + """汇总当前 Master 对应的容量、负载和缓存遥测。""" from lightllm.server.api_http import g_objs diff --git a/lightllm/server/httpserver_for_pd_master/admission.py b/lightllm/server/httpserver_for_pd_master/admission.py index 94ea62e53..9f25ebed8 100644 --- a/lightllm/server/httpserver_for_pd_master/admission.py +++ b/lightllm/server/httpserver_for_pd_master/admission.py @@ -38,6 +38,7 @@ class AdmissionPolicy: max_tracked_sessions: int = 100_000 def __post_init__(self) -> None: + """校验准入策略中的权重、超时和容量参数。""" if ( min( self.continuation_weight, @@ -66,6 +67,7 @@ def __post_init__(self) -> None: raise ValueError("max_tracked_sessions must be positive") def weight(self, priority: AdmissionPriority) -> int: + """返回指定优先级在轮转调度中的权重。""" if priority == AdmissionPriority.CONTINUATION: return self.continuation_weight if priority == AdmissionPriority.PROBABLE_CACHE_HIT: @@ -73,6 +75,7 @@ def weight(self, priority: AdmissionPriority) -> int: return self.cold_weight def max_wait_seconds(self, priority: AdmissionPriority) -> float: + """返回指定优先级允许的最长排队时间。""" if priority == AdmissionPriority.CONTINUATION: return self.continuation_max_wait_seconds if priority == AdmissionPriority.PROBABLE_CACHE_HIT: @@ -88,6 +91,7 @@ class AdmissionRequest: estimated_uncached_work: int = 0 def __post_init__(self) -> None: + """校验请求槽位数和预计未命中工作量。""" if self.decode_slots < 1: raise ValueError("decode_slots must be positive") if self.estimated_uncached_work < 0: @@ -102,11 +106,13 @@ class CacheCapacitySnapshot: capacity_tokens: int def __post_init__(self) -> None: + """校验缓存 token 统计值均为非负数。""" if self.total_tokens < 0 or self.capacity_tokens < 0: raise ValueError("cache token counts must be non-negative") @property def free_tokens(self) -> int: + """返回当前还能容纳的缓存 token 数。""" return max(0, self.capacity_tokens - self.total_tokens) @@ -119,12 +125,14 @@ def __init__( max_sessions: int, clock: Callable[[], float] = time.monotonic, ) -> None: + """初始化带 TTL 和数量上限的 Session 记录器。""" self.ttl_seconds = ttl_seconds self.max_sessions = max_sessions self._clock = clock self._last_success: OrderedDict[str, float] = OrderedDict() def is_continuation(self, session_key: Optional[str]) -> bool: + """判断 Session 是否在有效期内成功返回过结果。""" if not session_key: return False now = self._clock() @@ -138,6 +146,7 @@ def is_continuation(self, session_key: Optional[str]) -> bool: return True def mark_success(self, session_key: Optional[str]) -> None: + """记录 Session 最近一次成功返回结果的时间。""" if not session_key: return self._last_success[session_key] = self._clock() @@ -166,6 +175,7 @@ def __init__( waited_seconds: float, cold_slots: int, ) -> None: + """保存本次租约占用的总槽位和冷请求槽位。""" self._controller = controller self.request = request self.waited_seconds = waited_seconds @@ -173,12 +183,15 @@ def __init__( self._released = False async def __aenter__(self) -> "AdmissionLease": + """进入异步上下文并返回当前租约。""" return self async def __aexit__(self, _exc_type, _exc, _traceback) -> None: + """退出异步上下文时自动释放租约。""" self.release() def release(self) -> None: + """幂等释放本次占用的准入槽位。""" if self._released: return self._released = True @@ -202,6 +215,7 @@ def __init__( clock: Callable[[], float] = time.monotonic, state_change_callback: Optional[Callable[["PDAdmissionController"], None]] = None, ) -> None: + """初始化容量提供器、优先级队列和调度状态。""" self.policy = policy or AdmissionPolicy() self._decode_capacity_provider = decode_capacity_provider self._cache_capacity_provider = cache_capacity_provider @@ -223,10 +237,12 @@ def __init__( @property def active_slots(self) -> int: + """返回当前已经发放的 Decode 槽位数。""" return self._active_slots @property def active_cold_slots(self) -> int: + """返回当前由冷请求占用的槽位数。""" return self._active_cold_slots @property @@ -252,16 +268,20 @@ def cold_capacity(self) -> int: @property def queued_slots(self) -> int: + """返回等待队列中的 Decode 槽位总数。""" return self._queued_slots @property def queued_request_count(self) -> int: + """返回三个优先级队列中的请求总数。""" return sum(len(queue) for queue in self._queues.values()) def _capacity(self) -> int: + """读取并规范化当前可用的 Decode 容量。""" return max(0, int(self._decode_capacity_provider())) def _build_schedule(self) -> tuple[AdmissionPriority, ...]: + """按策略权重生成一个完整的轮转调度周期。""" remaining = {priority: self.policy.weight(priority) for priority in self._PRIORITY_ORDER} schedule = [] while any(remaining.values()): @@ -272,6 +292,7 @@ def _build_schedule(self) -> tuple[AdmissionPriority, ...]: return tuple(schedule) async def acquire(self, request: AdmissionRequest) -> AdmissionLease: + """立即发放租约或等待队列调度后再返回租约。""" capacity = self._capacity() if capacity <= 0 or request.decode_slots > capacity: raise ServerBusyError("PD decode capacity is unavailable") @@ -313,6 +334,7 @@ async def acquire(self, request: AdmissionRequest) -> AdmissionLease: raise def on_capacity_change(self) -> None: + """容量或缓存余量变化后重新尝试驱动队列。""" self._drain() def record_prefill_result( @@ -367,6 +389,7 @@ def promote_session(self, session_key: Optional[str]) -> None: self._drain() def _update_average(self, current: Optional[float], sample: float) -> float: + """按一个 Decode 波次大小更新指数移动平均。""" # 一个 Decode 波次作为自适应窗口:容量越大,单个样本对均值的影响越小。 sample_window = max(1, self._capacity()) alpha = 2.0 / (sample_window + 1.0) @@ -375,6 +398,7 @@ def _update_average(self, current: Optional[float], sample: float) -> float: return current + alpha * (sample - current) def _requires_cold_capacity(self, request: AdmissionRequest) -> bool: + """判断请求是否需要消耗冷请求容量。""" if request.priority == AdmissionPriority.COLD: return True return ( @@ -384,11 +408,13 @@ def _requires_cold_capacity(self, request: AdmissionRequest) -> bool: ) def _has_cold_capacity(self, request: AdmissionRequest) -> bool: + """判断剩余冷请求容量能否容纳当前请求。""" if not self._requires_cold_capacity(request): return True return self._active_cold_slots + request.decode_slots <= self.cold_capacity def _can_activate(self, request: AdmissionRequest) -> bool: + """检查总容量、冷容量和 Session 串行约束。""" if self._active_slots + request.decode_slots > self._capacity(): return False if not self._has_cold_capacity(request): @@ -396,6 +422,7 @@ def _can_activate(self, request: AdmissionRequest) -> bool: return request.session_key is None or request.session_key not in self._active_sessions def _activate(self, request: AdmissionRequest, waited_seconds: float) -> AdmissionLease: + """占用所需槽位并创建对应的准入租约。""" self._active_slots += request.decode_slots cold_slots = request.decode_slots if self._requires_cold_capacity(request) else 0 self._active_cold_slots += cold_slots @@ -404,6 +431,7 @@ def _activate(self, request: AdmissionRequest, waited_seconds: float) -> Admissi return AdmissionLease(self, request, waited_seconds, cold_slots) def _release(self, lease: AdmissionLease) -> None: + """归还租约槽位并继续调度等待请求。""" request = lease.request self._active_slots -= request.decode_slots self._active_cold_slots -= lease._cold_slots @@ -416,9 +444,11 @@ def _release(self, lease: AdmissionLease) -> None: self._drain() def _waiting_capacity(self) -> int: + """返回等待队列允许容纳的槽位总数。""" return self._capacity() * self.policy.waiting_decode_waves def _make_queue_room(self, incoming: _WaitingRequest) -> bool: + """必要时淘汰低优先级请求,为新请求腾出队列空间。""" waiting_capacity = self._waiting_capacity() if incoming.request.decode_slots > waiting_capacity: return False @@ -450,12 +480,14 @@ def _make_queue_room(self, incoming: _WaitingRequest) -> bool: return True def _enqueue(self, waiter: _WaitingRequest) -> None: + """把等待项加入优先级队列和 Session 队列。""" self._queues[waiter.request.priority].append(waiter) self._queued_slots += waiter.request.decode_slots if waiter.request.session_key is not None: self._session_queues.setdefault(waiter.request.session_key, deque()).append(waiter) def _remove_waiter(self, waiter: _WaitingRequest) -> bool: + """从所有索引中移除等待项并归还排队槽位。""" try: self._queues[waiter.request.priority].remove(waiter) except ValueError: @@ -471,6 +503,7 @@ def _remove_waiter(self, waiter: _WaitingRequest) -> bool: return True def _cancel_waiter_or_take_lease(self, waiter: _WaitingRequest) -> Optional[AdmissionLease]: + """取消等待项,或取回已并发发放的租约用于释放。""" if self._remove_waiter(waiter): waiter.future.cancel() self._drain() @@ -485,6 +518,7 @@ def _cancel_waiter_or_take_lease(self, waiter: _WaitingRequest) -> Optional[Admi return None async def _wait_for_lease(self, waiter: _WaitingRequest) -> AdmissionLease: + """等待租约、优先级提升后的新截止时间或超时。""" while True: deadline_changed_task = asyncio.create_task(waiter.deadline_changed.wait()) try: @@ -505,6 +539,7 @@ async def _wait_for_lease(self, waiter: _WaitingRequest) -> AdmissionLease: raise asyncio.TimeoutError def _session_is_grantable(self, waiter: _WaitingRequest) -> bool: + """判断等待项是否满足同 Session 串行和 FIFO 约束。""" session_key = waiter.request.session_key if session_key is None: return True @@ -513,6 +548,7 @@ def _session_is_grantable(self, waiter: _WaitingRequest) -> bool: return self._session_queues[session_key][0] is waiter def _first_grantable(self, priority: AdmissionPriority) -> Optional[_WaitingRequest]: + """返回指定优先级中当前最合适的可调度等待项。""" candidates = [] available_decode_slots = self._capacity() - self._active_slots for waiter in self._queues[priority]: @@ -538,6 +574,7 @@ def _first_grantable(self, priority: AdmissionPriority) -> Optional[_WaitingRequ ) def _first_fitting_higher_priority(self, blocked_priority: AdmissionPriority) -> Optional[_WaitingRequest]: + """查找能绕过受阻请求的更高优先级等待项。""" for priority in self._PRIORITY_ORDER: if priority <= blocked_priority: continue @@ -547,6 +584,7 @@ def _first_fitting_higher_priority(self, blocked_priority: AdmissionPriority) -> return None def _select_next(self) -> Optional[_WaitingRequest]: + """按照加权轮转顺序选择下一个等待项。""" schedule_size = len(self._schedule) for offset in range(schedule_size): index = (self._schedule_index + offset) % schedule_size @@ -557,6 +595,7 @@ def _select_next(self) -> Optional[_WaitingRequest]: return None def _drain(self) -> None: + """持续发放当前容量允许的等待请求。""" while self._active_slots < self._capacity(): schedule_index = self._schedule_index waiter = self._select_next() @@ -581,5 +620,6 @@ def _drain(self) -> None: self._notify_state_change() def _notify_state_change(self) -> None: + """通知外部记录最新的准入状态。""" if self._state_change_callback is not None: self._state_change_callback(self) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 6bf32108b..675ba3543 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -115,11 +115,13 @@ async def update_req_status(self, upkv_status: PDUpKVStatus): return def update_node_load_info(self, load_info: Optional[dict]) -> None: + """更新节点遥测并重新驱动准入队列。""" self.pd_manager.update_node_load_info(load_info) # Decode 租约或 Prefill cache 余量变化后都需要重新尝试队列。 self.admission_controller.on_capacity_change() def _record_admission_state(self, controller: PDAdmissionController) -> None: + """把当前准入队列状态写入监控指标。""" self.metric_client.gauge_set( "lightllm_pd_master_admission_queue_size", controller.queued_request_count, @@ -217,6 +219,7 @@ async def generate( admission_lease.release() def _get_session_key(self, request: Optional[Request]) -> Optional[str]: + """从请求头中提取规范化的 Session 标识。""" if request is None: return None session_key = request.headers.get("X-Session-Id", "").strip() @@ -228,6 +231,7 @@ def _build_admission_request( sampling_params: Optional[SamplingParams], session_key: Optional[str], ) -> AdmissionRequest: + """根据会话、缓存估算和 choice 数构造准入请求。""" estimated_cache_hit_rate = self.pd_manager.selector.estimate_prompt_cache_hit_rate(prompt) if estimated_cache_hit_rate is None or not math.isfinite(estimated_cache_hit_rate): estimated_cache_hit_rate = 0.0 @@ -868,6 +872,7 @@ def __init__(self, args: StartArgs): return def get_decode_capacity(self) -> int: + """汇总所有 Decode 节点租给当前 Master 的槽位。""" return sum( node.capacity_share if node.capacity_share is not None else node.start_args["running_max_req_size"] for node in self.decode_nodes diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index 4a5357c93..6bb9bdf6e 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -186,6 +186,7 @@ def estimate_cache_hit_rate(self, workers: List[PD_Client_Obj], request_text: st return min(max(result.matched_char_count / result.input_char_count, 0.0), 1.0) def _match_prompt_cache(self, request_text: str) -> PromptCacheMatchResult: + """优先复用当前请求保存的前缀匹配结果。""" match_context = _prompt_cache_match_context.get() # 准入和选点共享同一个提示词对象。通过对象身份判断可以避免比较可能很长的字符串, # ContextVar 则会将快照安全地传递给每个 n-choice 子任务。 diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py index 7bf844341..eb0f9b570 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py @@ -106,6 +106,7 @@ def record_prompt_cache_hit_rate(self, cache_hit_rate: float) -> None: self.policy.record_prompt_cache_hit_rate(cache_hit_rate) def estimate_prompt_cache_hit_rate(self, prompt: Union[str, List[int]]) -> Optional[float]: + """返回 cache-aware 策略对当前提示词的命中估算。""" if not isinstance(prompt, str): return 0.0 return self.policy.estimate_cache_hit_rate(self.prefill_nodes, prompt) From 81b4cfa29b81f0b833e053a531db857d54c581a3 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 1 Sep 2026 02:49:49 +0800 Subject: [PATCH 9/9] fix(pd): prevent admission collapse under saturation --- lightllm/server/api_cli.py | 2 +- lightllm/server/httpserver/pd_loop.py | 86 +--- .../httpserver_for_pd_master/admission.py | 477 +++++++++++------- .../httpserver_for_pd_master/manager.py | 210 ++++---- lightllm/server/metrics/metrics.py | 4 +- lightllm/server/pd_io_struct.py | 11 +- unit_tests/server/test_pd_admission.py | 448 +++++++++++----- unit_tests/server/test_pd_master_mode.py | 404 +++++++++++++-- 8 files changed, 1112 insertions(+), 530 deletions(-) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 16739babf..6500ad53b 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -71,7 +71,7 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: parser.add_argument( "--disable_pd_master_decode_capacity_limit", action="store_true", - help="Disable the PD master capacity and cache-aware admission queue.", + help="Disable the PD master admission queue based on registered decode capacity.", ) parser.add_argument( "--pd_trans_mode", diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index c3e3110b4..95a280bdf 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -12,24 +12,25 @@ import time from typing import Dict, Optional, Union, List from websockets import ClientConnection -from lightllm.server.pd_io_struct import NodeRole, ObjType +from lightllm.server.pd_io_struct import ( + NodeRole, + ObjType, + PD_MASTER_CAPACITY_EPOCH_KEY, + PD_MASTER_CAPACITY_SHARE_KEY, +) from lightllm.server.httpserver.async_queue import AsyncQueue from lightllm.utils.net_utils import get_hostname_ip from lightllm.utils.log_utils import init_logger -from lightllm.utils.envs_utils import get_lightllm_websocket_max_message_size, get_unique_server_name +from lightllm.utils.envs_utils import get_lightllm_websocket_max_message_size from lightllm.server.httpserver.manager import HttpServerManager from ..pd_io_struct import PD_Master_Obj from lightllm.server.core.objs import StartArgs from lightllm.server.core.objs import SamplingParams from lightllm.utils.error_utils import PDPrefillNodeStopGenToken from lightllm.utils.shm_port_args import get_shm_port_args -from lightllm.server.router.dynamic_prompt.radix_cache import RadixCacheReadOnlyClient logger = init_logger(__name__) -_radix_cache_client = None -_radix_cache_client_key = None - def _update_pd_master_membership(manager: HttpServerManager, pd_master_ids) -> None: """更新 Master 成员和容量版本,并立即唤醒心跳。""" @@ -56,40 +57,24 @@ def _allocate_capacity_share(total_capacity: int, pd_master_ids, pd_master_node_ return base + int(pd_master_ids.index(pd_master_node_id) < remainder) -def _get_radix_cache_info(): - """读取本节点各 DP 的 Radix cache token 统计。""" - global _radix_cache_client, _radix_cache_client_key - - from lightllm.server.api_http import g_objs - - args = g_objs.args - if args.disable_dynamic_prompt_cache: - return 0, 0, 0 - - max_total_token_num = g_objs.httpserver_manager.shm_max_total_token_num.get_value() - if max_total_token_num <= 0: - return 0, 0, 0 - - node_world_size = args.tp // args.nnodes - dp_world_size = args.tp // args.dp - client_key = (get_unique_server_name(), max_total_token_num, node_world_size, dp_world_size) - try: - if _radix_cache_client is None or _radix_cache_client_key != client_key: - _radix_cache_client = RadixCacheReadOnlyClient( - get_unique_server_name(), - max_total_token_num, - node_world_size=node_world_size, - dp_world_size=dp_world_size, - ) - _radix_cache_client_key = client_key - - dp_size_in_node = max(1, args.dp // args.nnodes) - total_tokens = sum(_radix_cache_client.get_tree_total_tokens_num(i) for i in range(dp_size_in_node)) - refed_tokens = sum(_radix_cache_client.get_refed_tokens_num(i) for i in range(dp_size_in_node)) - return int(total_tokens), int(refed_tokens), int(max_total_token_num * dp_size_in_node) - except Exception as exc: - logger.debug(f"read radix cache load failed: {str(exc)}") - return 0, 0, 0 +def _build_pd_registration_info(manager: HttpServerManager, pd_master_obj: PD_Master_Obj) -> dict: + """构造保持旧顶层 schema 兼容的 P/D 节点注册信息。""" + # Older Masters expand the registration JSON directly into PD_Client_Obj + # and reject unknown top-level fields during a rolling upgrade. + args_dict = vars(manager.args).copy() + args_dict["host"] = manager.host_ip + args_dict[PD_MASTER_CAPACITY_SHARE_KEY] = _allocate_capacity_share( + manager.args.running_max_req_size, + manager.pd_master_ids, + pd_master_obj.node_id, + ) + args_dict[PD_MASTER_CAPACITY_EPOCH_KEY] = manager.pd_master_capacity_epoch + return { + "node_id": manager.args.pd_node_id, + "client_ip_port": f"{manager.host_ip}:{get_shm_port_args().port}", + "mode": manager.pd_mode.value, + "start_args": args_dict, + } async def timer_log(manager: HttpServerManager): @@ -165,21 +150,8 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O sock = websocket.transport.get_extra_info("socket") sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) - args_dict = vars(manager.args) - args_dict["host"] = manager.host_ip # 发送注册信息 - regist_json = { - "node_id": manager.args.pd_node_id, - "client_ip_port": f"{manager.host_ip}:{get_shm_port_args().port}", - "mode": manager.pd_mode.value, - "start_args": args_dict, - "capacity_share": _allocate_capacity_share( - manager.args.running_max_req_size, - manager.pd_master_ids, - pd_master_obj.node_id, - ), - "capacity_epoch": manager.pd_master_capacity_epoch, - } + regist_json = _build_pd_registration_info(manager, pd_master_obj) await websocket.send(json.dumps(regist_json)) logger.info(f"Sent registration JSON: {regist_json}") @@ -375,7 +347,7 @@ async def _send_heartbeat_to_pd_master( # 获取节点负载信息 def _get_load_info(pd_master_node_id: int) -> dict: - """汇总当前 Master 对应的容量、负载和缓存遥测。""" + """汇总当前 Master 对应的容量和节点负载。""" from lightllm.server.api_http import g_objs @@ -389,15 +361,11 @@ def _get_load_info(pd_master_node_id: int) -> dict: float(g_objs.shared_token_load.get_dynamic_max_load(dp_index)) for dp_index in range(dp_size_in_node) ] mean_node_load = sum(current_load) / len(current_load) - radix_cache_total_tokens, radix_cache_refed_tokens, radix_cache_capacity_tokens = _get_radix_cache_info() pd_master_ids = getattr(g_objs.httpserver_manager, "pd_master_ids", (pd_master_node_id,)) load_info = { "total_token_usage_rate": mean_node_load, "client_ip_port": f"{g_objs.httpserver_manager.host_ip}:{get_shm_port_args().port}", "capacity_share": _allocate_capacity_share(args.running_max_req_size, pd_master_ids, pd_master_node_id), "capacity_epoch": getattr(g_objs.httpserver_manager, "pd_master_capacity_epoch", 0), - "radix_cache_total_tokens": radix_cache_total_tokens, - "radix_cache_refed_tokens": radix_cache_refed_tokens, - "radix_cache_capacity_tokens": radix_cache_capacity_tokens, } return load_info diff --git a/lightllm/server/httpserver_for_pd_master/admission.py b/lightllm/server/httpserver_for_pd_master/admission.py index 9f25ebed8..bebd02686 100644 --- a/lightllm/server/httpserver_for_pd_master/admission.py +++ b/lightllm/server/httpserver_for_pd_master/admission.py @@ -88,32 +88,11 @@ class AdmissionRequest: session_key: Optional[str] priority: AdmissionPriority decode_slots: int = 1 - estimated_uncached_work: int = 0 def __post_init__(self) -> None: - """校验请求槽位数和预计未命中工作量。""" + """校验请求需要原子获取的 Decode 槽位数。""" if self.decode_slots < 1: raise ValueError("decode_slots must be positive") - if self.estimated_uncached_work < 0: - raise ValueError("estimated_uncached_work must be non-negative") - - -@dataclass(frozen=True, slots=True) -class CacheCapacitySnapshot: - """当前 PD Master 可使用的 Prefill Radix cache 份额。""" - - total_tokens: int - capacity_tokens: int - - def __post_init__(self) -> None: - """校验缓存 token 统计值均为非负数。""" - if self.total_tokens < 0 or self.capacity_tokens < 0: - raise ValueError("cache token counts must be non-negative") - - @property - def free_tokens(self) -> int: - """返回当前还能容纳的缓存 token 数。""" - return max(0, self.capacity_tokens - self.total_tokens) class SessionTracker: @@ -157,7 +136,6 @@ def mark_success(self, session_key: Optional[str]) -> None: @dataclass(slots=True) class _WaitingRequest: - sequence_id: int request: AdmissionRequest enqueue_time: float deadline: float @@ -173,13 +151,11 @@ def __init__( controller: "PDAdmissionController", request: AdmissionRequest, waited_seconds: float, - cold_slots: int, ) -> None: - """保存本次租约占用的总槽位和冷请求槽位。""" + """保存本次租约占用的 Decode 槽位。""" self._controller = controller self.request = request self.waited_seconds = waited_seconds - self._cold_slots = cold_slots self._released = False async def __aenter__(self) -> "AdmissionLease": @@ -210,7 +186,6 @@ class PDAdmissionController: def __init__( self, decode_capacity_provider: Callable[[], int], - cache_capacity_provider: Optional[Callable[[], Optional[CacheCapacitySnapshot]]] = None, policy: Optional[AdmissionPolicy] = None, clock: Callable[[], float] = time.monotonic, state_change_callback: Optional[Callable[["PDAdmissionController"], None]] = None, @@ -218,54 +193,28 @@ def __init__( """初始化容量提供器、优先级队列和调度状态。""" self.policy = policy or AdmissionPolicy() self._decode_capacity_provider = decode_capacity_provider - self._cache_capacity_provider = cache_capacity_provider self._clock = clock self._state_change_callback = state_change_callback self._active_slots = 0 - self._active_cold_slots = 0 - self._average_cold_uncached_tokens: Optional[float] = None - self._probable_actual_hit_rate: Optional[float] = None self._active_sessions = set() self._queues: Dict[AdmissionPriority, Deque[_WaitingRequest]] = { priority: deque() for priority in self._PRIORITY_ORDER } self._session_queues: Dict[str, Deque[_WaitingRequest]] = {} self._queued_slots = 0 - self._sequence_id = 0 - self._schedule = self._build_schedule() - self._schedule_index = 0 + self._deficits: Dict[AdmissionPriority, int] = {priority: 0 for priority in self._PRIORITY_ORDER} + self._priority_index = 0 + self._priority_visit_started = False + self._blocked_waiter: Optional[_WaitingRequest] = None + self._backfilled_slots = 0 + self._backfill_limit = 0 + self._reservation_active = False @property def active_slots(self) -> int: """返回当前已经发放的 Decode 槽位数。""" return self._active_slots - @property - def active_cold_slots(self) -> int: - """返回当前由冷请求占用的槽位数。""" - return self._active_cold_slots - - @property - def cold_capacity(self) -> int: - """返回在当前缓存余量下允许并发的冷请求槽位数。""" - decode_capacity = self._capacity() - if ( - decode_capacity <= 0 - or self._cache_capacity_provider is None - or self._average_cold_uncached_tokens is None - or self._average_cold_uncached_tokens <= 0 - ): - return decode_capacity - - snapshot = self._cache_capacity_provider() - if snapshot is None or snapshot.capacity_tokens <= 0: - return decode_capacity - - # 剩余缓存能容纳几个“平均冷请求”,就开放几个冷槽位;至少保留一个 - # 探索槽位,使系统在缓存已满时仍能接纳新会话并持续获得反馈。 - requests_fitting_in_cache = int(snapshot.free_tokens / self._average_cold_uncached_tokens) - return min(decode_capacity, max(1, requests_fitting_in_cache)) - @property def queued_slots(self) -> int: """返回等待队列中的 Decode 槽位总数。""" @@ -280,39 +229,31 @@ def _capacity(self) -> int: """读取并规范化当前可用的 Decode 容量。""" return max(0, int(self._decode_capacity_provider())) - def _build_schedule(self) -> tuple[AdmissionPriority, ...]: - """按策略权重生成一个完整的轮转调度周期。""" - remaining = {priority: self.policy.weight(priority) for priority in self._PRIORITY_ORDER} - schedule = [] - while any(remaining.values()): - for priority in self._PRIORITY_ORDER: - if remaining[priority] > 0: - schedule.append(priority) - remaining[priority] -= 1 - return tuple(schedule) - async def acquire(self, request: AdmissionRequest) -> AdmissionLease: """立即发放租约或等待队列调度后再返回租约。""" capacity = self._capacity() if capacity <= 0 or request.decode_slots > capacity: raise ServerBusyError("PD decode capacity is unavailable") - if self.queued_request_count == 0 and self._can_activate(request): + if self.queued_request_count == 0 and self._can_activate(request, capacity): lease = self._activate(request, waited_seconds=0.0) self._notify_state_change() return lease + idle_fill_lease = self._try_activate_idle_fill(request, capacity) + if idle_fill_lease is not None: + self._notify_state_change() + return idle_fill_lease + loop = asyncio.get_running_loop() enqueue_time = self._clock() waiter = _WaitingRequest( - sequence_id=self._sequence_id, request=request, enqueue_time=enqueue_time, deadline=enqueue_time + self.policy.max_wait_seconds(request.priority), deadline_changed=asyncio.Event(), future=loop.create_future(), ) - self._sequence_id += 1 if not self._make_queue_room(waiter): raise ServerBusyError("PD master admission queue is full") @@ -334,36 +275,9 @@ async def acquire(self, request: AdmissionRequest) -> AdmissionLease: raise def on_capacity_change(self) -> None: - """容量或缓存余量变化后重新尝试驱动队列。""" - self._drain() - - def record_prefill_result( - self, - request: AdmissionRequest, - prompt_tokens: int, - cached_tokens: int, - ) -> None: - """用真实命中结果更新预计命中的可信度和冷请求容量。""" - if request.priority == AdmissionPriority.CONTINUATION: - return - - prompt_tokens = max(0, int(prompt_tokens)) - cached_tokens = min(prompt_tokens, max(0, int(cached_tokens))) - uncached_tokens = prompt_tokens - cached_tokens - - if request.priority == AdmissionPriority.PROBABLE_CACHE_HIT: - actual_hit_rate = cached_tokens / max(prompt_tokens, 1) - self._probable_actual_hit_rate = self._update_average( - self._probable_actual_hit_rate, - actual_hit_rate, - ) - - # 预计命中的真实命中率低于承诺阈值时,让后续同类请求也消费冷槽位。 - if self._requires_cold_capacity(request): - self._average_cold_uncached_tokens = self._update_average( - self._average_cold_uncached_tokens, - float(uncached_tokens), - ) + """Decode 容量变化后重置临时公平状态并重新驱动队列。""" + self._clear_backfill_state() + self._reset_deficits() self._drain() def promote_session(self, session_key: Optional[str]) -> None: @@ -374,6 +288,8 @@ def promote_session(self, session_key: Optional[str]) -> None: if not session_queue: return + if self._blocked_waiter in session_queue: + self._clear_backfill_state(reset_deficits=True) for waiter in tuple(session_queue): old_priority = waiter.request.priority if old_priority == AdmissionPriority.CONTINUATION: @@ -386,59 +302,76 @@ def promote_session(self, session_key: Optional[str]) -> None: ) waiter.deadline_changed.set() self._queues[AdmissionPriority.CONTINUATION].append(waiter) + if not self._queues[old_priority]: + self._deficits[old_priority] = 0 self._drain() - def _update_average(self, current: Optional[float], sample: float) -> float: - """按一个 Decode 波次大小更新指数移动平均。""" - # 一个 Decode 波次作为自适应窗口:容量越大,单个样本对均值的影响越小。 - sample_window = max(1, self._capacity()) - alpha = 2.0 / (sample_window + 1.0) - if current is None: - return sample - return current + alpha * (sample - current) - - def _requires_cold_capacity(self, request: AdmissionRequest) -> bool: - """判断请求是否需要消耗冷请求容量。""" - if request.priority == AdmissionPriority.COLD: - return True - return ( - request.priority == AdmissionPriority.PROBABLE_CACHE_HIT - and self._probable_actual_hit_rate is not None - and self._probable_actual_hit_rate < self.policy.probable_cache_hit_threshold - ) - - def _has_cold_capacity(self, request: AdmissionRequest) -> bool: - """判断剩余冷请求容量能否容纳当前请求。""" - if not self._requires_cold_capacity(request): - return True - return self._active_cold_slots + request.decode_slots <= self.cold_capacity - - def _can_activate(self, request: AdmissionRequest) -> bool: - """检查总容量、冷容量和 Session 串行约束。""" - if self._active_slots + request.decode_slots > self._capacity(): - return False - if not self._has_cold_capacity(request): + def _can_activate(self, request: AdmissionRequest, capacity: Optional[int] = None) -> bool: + """检查 Decode 总容量和 Session 串行约束。""" + if capacity is None: + capacity = self._capacity() + if self._active_slots + request.decode_slots > capacity: return False return request.session_key is None or request.session_key not in self._active_sessions + def _try_activate_idle_fill( + self, + request: AdmissionRequest, + capacity: int, + ) -> Optional[AdmissionLease]: + """在满队列拒绝前,用当前唯一可运行的新请求填充空槽。 + + 只覆盖两种不会越过可运行旧请求的场景:现有队列全部受 + Session 串行约束,或已保护 gang 仍在一波 bounded backfill 预算内。 + """ + if not self._can_activate(request, capacity): + return None + if request.session_key is not None and request.session_key in self._session_queues: + return None + + blocked = self._blocked_waiter + if blocked is None: + has_grantable_waiter = any( + not waiter.future.done() and self._session_is_grantable(waiter) + for queue in self._queues.values() + for waiter in queue + ) + if has_grantable_waiter: + return None + return self._activate(request, waited_seconds=0.0) + + available_slots = capacity - self._active_slots + if ( + self._reservation_active + or blocked.future.done() + or not self._session_is_grantable(blocked) + or self._has_fitting_backfill(blocked, available_slots) + ): + return None + + remaining_backfill = self._backfill_limit - self._backfilled_slots + if request.decode_slots > remaining_backfill: + return None + + lease = self._activate(request, waited_seconds=0.0) + self._backfilled_slots += request.decode_slots + if self._backfilled_slots >= self._backfill_limit: + self._reservation_active = True + return lease + def _activate(self, request: AdmissionRequest, waited_seconds: float) -> AdmissionLease: """占用所需槽位并创建对应的准入租约。""" self._active_slots += request.decode_slots - cold_slots = request.decode_slots if self._requires_cold_capacity(request) else 0 - self._active_cold_slots += cold_slots if request.session_key is not None: self._active_sessions.add(request.session_key) - return AdmissionLease(self, request, waited_seconds, cold_slots) + return AdmissionLease(self, request, waited_seconds) def _release(self, lease: AdmissionLease) -> None: """归还租约槽位并继续调度等待请求。""" request = lease.request self._active_slots -= request.decode_slots - self._active_cold_slots -= lease._cold_slots if self._active_slots < 0: raise RuntimeError("PD admission active slot count became negative") - if self._active_cold_slots < 0: - raise RuntimeError("PD admission active cold slot count became negative") if request.session_key is not None: self._active_sessions.discard(request.session_key) self._drain() @@ -500,6 +433,12 @@ def _remove_waiter(self, waiter: _WaitingRequest) -> bool: session_queue.remove(waiter) if not session_queue: self._session_queues.pop(session_key, None) + if waiter is self._blocked_waiter: + # 非正常移除(取消、超时、替换、缩容)放弃已经预扣的 gang + # 服务机会;重置 DRR 状态比跨优先级退款更安全。 + self._clear_backfill_state(reset_deficits=True) + if not self._queues[waiter.request.priority]: + self._deficits[waiter.request.priority] = 0 return True def _cancel_waiter_or_take_lease(self, waiter: _WaitingRequest) -> Optional[AdmissionLease]: @@ -547,76 +486,228 @@ def _session_is_grantable(self, waiter: _WaitingRequest) -> bool: return False return self._session_queues[session_key][0] is waiter - def _first_grantable(self, priority: AdmissionPriority) -> Optional[_WaitingRequest]: - """返回指定优先级中当前最合适的可调度等待项。""" - candidates = [] - available_decode_slots = self._capacity() - self._active_slots + def _first_grantable( + self, + priority: AdmissionPriority, + excluded: Optional[_WaitingRequest] = None, + available_slots: Optional[int] = None, + ) -> Optional[_WaitingRequest]: + """按类内 FIFO 返回第一个满足 Session 和可选槽位约束的等待项。""" for waiter in self._queues[priority]: - if not self._session_is_grantable(waiter): + if waiter is excluded: continue - if priority != AdmissionPriority.COLD: - return waiter - if waiter.request.decode_slots > self.cold_capacity: + if not self._session_is_grantable(waiter): continue - if waiter.request.decode_slots <= available_decode_slots and not self._has_cold_capacity(waiter.request): + if available_slots is not None and waiter.request.decode_slots > available_slots: continue - candidates.append(waiter) + return waiter + return None + + def _advance_priority(self) -> None: + """结束当前 DRR 类访问并移到下一优先级。""" + self._priority_index = (self._priority_index + 1) % len(self._PRIORITY_ORDER) + self._priority_visit_started = False + def _reset_deficits(self) -> None: + """清空按 Decode 槽位计费的 DRR 临时信用。""" + for priority in self._PRIORITY_ORDER: + self._deficits[priority] = 0 + self._priority_index = 0 + self._priority_visit_started = False + + def _select_weighted( + self, + capacity: int, + excluded: Optional[_WaitingRequest] = None, + available_slots: Optional[int] = None, + ) -> Optional[_WaitingRequest]: + """用按槽位计费的 deficit round-robin 选择一个等待项。""" + candidates = { + priority: waiter + for priority in self._PRIORITY_ORDER + if ( + waiter := self._first_grantable( + priority, + excluded=excluded, + available_slots=available_slots, + ) + ) + is not None + } + for priority in self._PRIORITY_ORDER: + # 空类和暂时全部受 Session 串行约束的类都不能积攒无限信用。 + if self._first_grantable(priority) is None: + self._deficits[priority] = 0 if not candidates: return None - # 冷请求内部优先处理预计新增缓存最少的任务,在同等代价下保持 FIFO。 - return min( - candidates, - key=lambda waiter: ( - waiter.request.estimated_uncached_work, - waiter.sequence_id, - ), - ) - def _first_fitting_higher_priority(self, blocked_priority: AdmissionPriority) -> Optional[_WaitingRequest]: - """查找能绕过受阻请求的更高优先级等待项。""" - for priority in self._PRIORITY_ORDER: - if priority <= blocked_priority: + only_priority = next(iter(candidates)) if len(candidates) == 1 else None + while True: + priority = self._PRIORITY_ORDER[self._priority_index] + waiter = candidates.get(priority) + if waiter is None: + self._advance_priority() continue - for waiter in self._queues[priority]: - if self._session_is_grantable(waiter) and self._can_activate(waiter.request): - return waiter - return None - def _select_next(self) -> Optional[_WaitingRequest]: - """按照加权轮转顺序选择下一个等待项。""" - schedule_size = len(self._schedule) - for offset in range(schedule_size): - index = (self._schedule_index + offset) % schedule_size - waiter = self._first_grantable(self._schedule[index]) - if waiter is not None: - self._schedule_index = (index + 1) % schedule_size + if not self._priority_visit_started: + self._deficits[priority] += self.policy.weight(priority) + self._priority_visit_started = True + + # 单一活跃类必须保持 work-conserving;直接补足若干轮 quantum, + # 避免大 gang 仅因 DRR 信用暂时不足而留下 Decode 空槽。 + if only_priority == priority and self._deficits[priority] < waiter.request.decode_slots: + quantum = self.policy.weight(priority) + missing = waiter.request.decode_slots - self._deficits[priority] + visits = (missing + quantum - 1) // quantum + self._deficits[priority] += visits * quantum + + if waiter.request.decode_slots <= self._deficits[priority]: + # 在选择点统一按 choice 槽位扣费。即使 gang 暂时因物理空槽 + # 不足进入 backfill,它的 DRR 服务机会也已经被完整计费。 + self._deficits[priority] -= waiter.request.decode_slots return waiter - return None + self._advance_priority() + + def _clear_backfill_state(self, reset_deficits: bool = False) -> None: + """清除 gang backfill 或 reservation 的全部临时状态。""" + self._blocked_waiter = None + self._backfilled_slots = 0 + self._backfill_limit = 0 + self._reservation_active = False + if reset_deficits: + self._reset_deficits() + + def _start_backfill(self, waiter: _WaitingRequest, capacity: int) -> None: + """为仅受当前可用槽位阻塞的 gang 启动一波有限 backfill。""" + self._blocked_waiter = waiter + self._backfilled_slots = 0 + self._backfill_limit = capacity + self._reservation_active = False + + def _fail_oversized_waiters(self, capacity: int) -> None: + """容量缩小时失败掉已经不可能原子获得所需槽位的等待项。""" + for priority in self._PRIORITY_ORDER: + for waiter in tuple(self._queues[priority]): + if waiter.request.decode_slots <= capacity: + continue + if self._remove_waiter(waiter) and not waiter.future.done(): + waiter.future.set_exception(ServerBusyError("PD decode capacity fell below queued request size")) + + def _trim_queue_to_capacity(self, capacity: int) -> None: + """容量缩小时按低优先级、同级最新顺序恢复等待队列上限。""" + waiting_capacity = capacity * self.policy.waiting_decode_waves + slots_to_remove = self._queued_slots - waiting_capacity + if slots_to_remove <= 0: + return - def _drain(self) -> None: - """持续发放当前容量允许的等待请求。""" - while self._active_slots < self._capacity(): - schedule_index = self._schedule_index - waiter = self._select_next() - if waiter is None: + victims = [] + removed_slots = 0 + for priority in reversed(self._PRIORITY_ORDER): + for waiter in reversed(self._queues[priority]): + victims.append(waiter) + removed_slots += waiter.request.decode_slots + if removed_slots >= slots_to_remove: + break + if removed_slots >= slots_to_remove: break - if not self._can_activate(waiter.request): - # 为被选中的多 choice 请求积累槽位,但不因此阻塞当前可以执行的更高优先级请求。 - self._schedule_index = schedule_index - waiter = self._first_fitting_higher_priority(waiter.request.priority) + + for victim in victims: + if self._remove_waiter(victim) and not victim.future.done(): + victim.future.set_exception(ServerBusyError("PD master admission queue capacity shrank")) + + def _grant_waiter(self, waiter: _WaitingRequest) -> bool: + """从队列移除已由 DRR 计费的等待项并原子发放租约。""" + if waiter.future.done(): + self._remove_waiter(waiter) + self._reset_deficits() + return False + + priority = waiter.request.priority + if waiter is self._blocked_waiter: + # 正常兑现 reservation 时保留选择点已经完成的 DRR 扣费。 + self._clear_backfill_state() + if not self._remove_waiter(waiter): + self._reset_deficits() + return False + if not self._queues[priority]: + self._deficits[priority] = 0 + + lease = self._activate( + waiter.request, + waited_seconds=max(0.0, self._clock() - waiter.enqueue_time), + ) + waiter.future.set_result(lease) + return True + + def _has_fitting_backfill(self, blocked: _WaitingRequest, available_slots: int) -> bool: + """判断是否有不含被保护 gang 的请求当前可以填充空槽。""" + return any( + self._first_grantable( + priority, + excluded=blocked, + available_slots=available_slots, + ) + is not None + for priority in self._PRIORITY_ORDER + ) + + def _drain(self) -> None: + """持续发放租约,并为受空槽碎片阻塞的 gang 提供有限 backfill。""" + capacity = self._capacity() + self._fail_oversized_waiters(capacity) + self._trim_queue_to_capacity(capacity) + + while self._active_slots < capacity: + available_slots = capacity - self._active_slots + blocked = self._blocked_waiter + if blocked is not None: + if blocked.future.done(): + self._remove_waiter(blocked) + continue + if not self._session_is_grantable(blocked): + # Session 阻塞不是容量碎片,不能借此获得全局 reservation。 + self._clear_backfill_state(reset_deficits=True) + continue + if blocked.request.decode_slots <= available_slots: + self._grant_waiter(blocked) + continue + if self._reservation_active: + break + + remaining_backfill = self._backfill_limit - self._backfilled_slots + if remaining_backfill <= 0: + self._reservation_active = True + break + backfill_slots = min(available_slots, remaining_backfill) + waiter = self._select_weighted( + capacity, + excluded=blocked, + available_slots=backfill_slots, + ) if waiter is None: + # 没有可填当前空槽的请求时保留 backfill 机会;稍后到达的小请求 + # 仍可使用本波预算。只有存在物理上可填、但会越过预算的请求时 + # 才立即转入 reservation。 + if self._has_fitting_backfill(blocked, available_slots): + self._reservation_active = True break - if not self._remove_waiter(waiter): + granted_slots = waiter.request.decode_slots + if not self._grant_waiter(waiter): + continue + self._backfilled_slots += granted_slots + if self._backfilled_slots >= self._backfill_limit: + self._reservation_active = True continue - lease = self._activate( - waiter.request, - waited_seconds=max(0.0, self._clock() - waiter.enqueue_time), - ) - if waiter.future.done(): - lease.release() + + waiter = self._select_weighted(capacity) + if waiter is None: + break + if waiter.request.decode_slots > available_slots: + # 此处 waiter 已满足 Session 约束且不超过总容量,唯一阻塞原因 + # 是当前空槽不足,因此可以安全启动 bounded backfill。 + self._start_backfill(waiter, capacity) continue - waiter.future.set_result(lease) + self._grant_waiter(waiter) self._notify_state_change() def _notify_state_change(self) -> None: diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 675ba3543..fba20821d 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -13,7 +13,14 @@ asyncio.set_event_loop_policy(uvloop.EventLoopPolicy()) from typing import Union, List, Tuple, Dict, Optional from lightllm.server.core.objs import FinishStatus -from ..pd_io_struct import PD_Client_Obj, PDUpKVStatus, ObjType, PDDecodeNodeInfo +from ..pd_io_struct import ( + PD_Client_Obj, + PDUpKVStatus, + ObjType, + PDDecodeNodeInfo, + PD_MASTER_CAPACITY_EPOCH_KEY, + PD_MASTER_CAPACITY_SHARE_KEY, +) from lightllm.server.core.objs import SamplingParams, StartArgs from ..multimodal_params import MultimodalParams from ..tokenizer import get_tokenizer @@ -31,7 +38,6 @@ AdmissionPolicy, AdmissionPriority, AdmissionRequest, - CacheCapacitySnapshot, PDAdmissionController, SessionTracker, ) @@ -57,9 +63,11 @@ def __init__( ttl_seconds=self.admission_policy.active_session_ttl_seconds, max_sessions=self.admission_policy.max_tracked_sessions, ) + self._last_admission_metric_values: Dict[str, int] = {} + self._pending_admission_metric_values: Optional[Dict[str, int]] = None + self._admission_metric_flush_scheduled = False self.admission_controller = PDAdmissionController( decode_capacity_provider=self.pd_manager.get_decode_capacity, - cache_capacity_provider=self.pd_manager.get_prefill_cache_capacity, policy=self.admission_policy, state_change_callback=self._record_admission_state, ) @@ -115,25 +123,42 @@ async def update_req_status(self, upkv_status: PDUpKVStatus): return def update_node_load_info(self, load_info: Optional[dict]) -> None: - """更新节点遥测并重新驱动准入队列。""" - self.pd_manager.update_node_load_info(load_info) - # Decode 租约或 Prefill cache 余量变化后都需要重新尝试队列。 - self.admission_controller.on_capacity_change() + """更新节点遥测;仅 Decode 容量租约变化时重新驱动准入队列。""" + if self.pd_manager.update_node_load_info(load_info): + self.admission_controller.on_capacity_change() def _record_admission_state(self, controller: PDAdmissionController) -> None: - """把当前准入队列状态写入监控指标。""" - self.metric_client.gauge_set( - "lightllm_pd_master_admission_queue_size", - controller.queued_request_count, - ) - self.metric_client.gauge_set( - "lightllm_pd_master_admission_active_slots", - controller.active_slots, - ) - self.metric_client.gauge_set( - "lightllm_pd_master_admission_cold_capacity", - controller.cold_capacity, - ) + """合并同一事件循环周期内的状态变化,避免重复发送 gauge RPC。""" + self._pending_admission_metric_values = { + "lightllm_pd_master_admission_queue_size": controller.queued_request_count, + "lightllm_pd_master_admission_active_slots": controller.active_slots, + "lightllm_pd_master_admission_decode_capacity": self.pd_manager.get_decode_capacity(), + } + if self._admission_metric_flush_scheduled: + return + + try: + loop = asyncio.get_running_loop() + except RuntimeError: + self._flush_admission_state_metrics() + return + + self._admission_metric_flush_scheduled = True + loop.call_soon(self._flush_admission_state_metrics) + + def _flush_admission_state_metrics(self) -> None: + """只发送相较于上次上报实际发生变化的 admission gauge。""" + self._admission_metric_flush_scheduled = False + metric_values = self._pending_admission_metric_values + self._pending_admission_metric_values = None + if metric_values is None: + return + + for name, value in metric_values.items(): + if self._last_admission_metric_values.get(name) == value: + continue + self.metric_client.gauge_set(name, value) + self._last_admission_metric_values[name] = value def tokens(self, prompt, multimodal_params, samping_params: SamplingParams, kwargs=None): kwargs = {} if kwargs is None else kwargs @@ -174,47 +199,33 @@ async def generate( request: Request, ): admission_lease = None - admission_request = None - observed_prefill_ids = set() + running_request_registered = False session_key = self._get_session_key(request) - if not self.args.disable_pd_master_decode_capacity_limit: - admission_request = self._build_admission_request(prompt, sampling_params, session_key) - admission_lease = await self.admission_controller.acquire(admission_request) - admission_request = admission_lease.request - self.metric_client.histogram_observe( - "lightllm_request_queue_duration_bucket", admission_lease.waited_seconds - ) - if admission_lease.waited_seconds > 0: - # 排队期间节点和缓存内容可能变化;派发前重新匹配,避免复用过期快照。 - self.pd_manager.selector.estimate_prompt_cache_hit_rate(prompt) - - was_idle = self.running_request_count == 0 - self.running_request_count += 1 - if was_idle: - self.latest_success_infer_time = time.time() try: + if not self.args.disable_pd_master_decode_capacity_limit: + admission_request = self._build_admission_request(prompt, sampling_params, session_key) + admission_lease = await self.admission_controller.acquire(admission_request) + self.metric_client.histogram_observe( + "lightllm_request_queue_duration_bucket", admission_lease.waited_seconds + ) + if admission_lease.waited_seconds > 0: + # 排队期间节点和缓存内容可能变化;派发前重新匹配,避免复用过期快照。 + self.pd_manager.selector.estimate_prompt_cache_hit_rate(prompt) + + was_idle = self.running_request_count == 0 + self.running_request_count += 1 + running_request_registered = True + if was_idle: + self.latest_success_infer_time = time.time() async with aclosing(self._generate(prompt, sampling_params, multimodal_params, request)) as generator: async for result in generator: - if ( - admission_request is not None - and isinstance(result, tuple) - and len(result) >= 3 - and result[0] not in observed_prefill_ids - and isinstance(result[2], dict) - and "prompt_tokens" in result[2] - ): - observed_prefill_ids.add(result[0]) - self.admission_controller.record_prefill_result( - admission_request, - result[2]["prompt_tokens"], - result[2].get("prompt_cache_len", 0), - ) if session_key is not None and not self.session_tracker.is_continuation(session_key): self.session_tracker.mark_success(session_key) self.admission_controller.promote_session(session_key) yield result finally: - self.running_request_count -= 1 + if running_request_registered: + self.running_request_count -= 1 if admission_lease is not None: admission_lease.release() @@ -245,13 +256,10 @@ def _build_admission_request( priority = AdmissionPriority.COLD decode_slots = max(1, int(getattr(sampling_params, "n", 1) or 1)) - prompt_size = len(prompt) if prompt is not None else 0 - estimated_uncached_work = math.ceil(prompt_size * (1.0 - estimated_cache_hit_rate)) * decode_slots return AdmissionRequest( session_key=session_key, priority=priority, decode_slots=decode_slots, - estimated_uncached_work=estimated_uncached_work, ) async def _generate( @@ -878,44 +886,6 @@ def get_decode_capacity(self) -> int: for node in self.decode_nodes ) - def get_prefill_cache_capacity(self) -> Optional[CacheCapacitySnapshot]: - """汇总当前 Master 对应的 Prefill cache 份额;遥测不完整时不参与限流。""" - if not self.prefill_nodes: - return None - - statuses = [node.run_status for node in self.prefill_nodes] - if any(status.radix_cache_capacity_tokens <= 0 or status.report_time <= 0 for status in statuses): - return None - - full_decode_capacity = sum(node.start_args["running_max_req_size"] for node in self.decode_nodes) - local_decode_capacity = self.get_decode_capacity() - if full_decode_capacity <= 0 or local_decode_capacity <= 0: - return None - - # 所有 Master 都能看到同一组 P 节点,因此按本 Master 的 Decode 租约比例 - # 切分缓存余量,避免每个 Master 重复消费整份 headroom。 - share_ratio = min(1.0, local_decode_capacity / full_decode_capacity) - # total_token_usage_rate 已排除可驱逐 Radix token,因此下面两项分别表示 - # 当前可驱逐缓存量和运行中请求之外还能留给缓存的容量。 - total_tokens = int( - sum(max(0, status.radix_cache_total_tokens - status.radix_cache_refed_tokens) for status in statuses) - * share_ratio - ) - capacity_tokens = int( - sum( - max( - 0, - status.radix_cache_capacity_tokens * (1.0 - min(max(status.total_token_usage_rate, 0.0), 1.0)), - ) - for status in statuses - ) - * share_ratio - ) - return CacheCapacitySnapshot( - total_tokens=total_tokens, - capacity_tokens=capacity_tokens, - ) - def is_pd_nodes_ready(self): prefill_node_count = len(self.prefill_nodes) decode_node_count = len(self.decode_nodes) @@ -967,7 +937,24 @@ async def check_pd_nodes_health(self): return True def register_pd(self, pd_info_json, websocket): - pd_client = PD_Client_Obj(**pd_info_json) + # Capacity metadata lives in reserved start_args keys so newer P/D + # nodes retain the legacy registration schema for older Masters. Keep + # accepting the short-lived top-level form for branch compatibility. + pd_info = dict(pd_info_json) + start_args = pd_info.get("start_args") or {} + capacity_share = pd_info.pop( + "capacity_share", + start_args.get(PD_MASTER_CAPACITY_SHARE_KEY), + ) + capacity_epoch = pd_info.pop( + "capacity_epoch", + start_args.get(PD_MASTER_CAPACITY_EPOCH_KEY, 0), + ) + pd_client = PD_Client_Obj( + **pd_info, + capacity_share=capacity_share, + capacity_epoch=capacity_epoch, + ) client_max_req_total_len = pd_client.start_args["max_req_total_len"] if client_max_req_total_len != self.args.max_req_total_len: logger.error( @@ -1007,19 +994,25 @@ def register_pd(self, pd_info_json, websocket): return def remove_pd(self, pd_info_json): - pd_client = PD_Client_Obj(**pd_info_json) + # Disconnect cleanup only needs the stable legacy identity fields; do + # not reconstruct PD_Client_Obj from a possibly newer registration. + client_ip_port = pd_info_json["client_ip_port"] + mode = pd_info_json.get("mode", "unknown") - self.url_to_pd_nodes.pop(pd_client.client_ip_port, None) - self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != pd_client.client_ip_port] - self.decode_nodes = [e for e in self.decode_nodes if e.client_ip_port != pd_client.client_ip_port] + self.url_to_pd_nodes.pop(client_ip_port, None) + self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != client_ip_port] + self.decode_nodes = [e for e in self.decode_nodes if e.client_ip_port != client_ip_port] self.selector.update_nodes(self.prefill_nodes, self.decode_nodes) - logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} removed") + logger.info(f"mode: {mode} url: {client_ip_port} removed") return - def update_node_load_info(self, load_info: Optional[dict]): - """更新节点负载信息 + def update_node_load_info(self, load_info: Optional[dict]) -> bool: + """更新节点负载信息,并返回有效 Decode 容量份额是否发生变化。 + + capacity_epoch 只用于拒绝旧上报;仅 epoch 变化不会重新驱动 admission。 + load_info: 节点负载信息字典,内容格式如下,可以为 None { "total_token_usage_rate": xxxx, @@ -1028,16 +1021,12 @@ def update_node_load_info(self, load_info: Optional[dict]): """ try: if load_info is None: - return + return False client_ip_port = load_info["client_ip_port"] pd_client = self.url_to_pd_nodes.get(client_ip_port) if pd_client is None: - return + return False pd_client.run_status.total_token_usage_rate = load_info["total_token_usage_rate"] - pd_client.run_status.radix_cache_total_tokens = load_info.get("radix_cache_total_tokens", 0) - pd_client.run_status.radix_cache_refed_tokens = load_info.get("radix_cache_refed_tokens", 0) - pd_client.run_status.radix_cache_capacity_tokens = load_info.get("radix_cache_capacity_tokens", 0) - pd_client.run_status.report_time = time.monotonic() capacity_epoch = int(load_info.get("capacity_epoch", pd_client.capacity_epoch)) if capacity_epoch >= pd_client.capacity_epoch: @@ -1046,11 +1035,14 @@ def update_node_load_info(self, load_info: Optional[dict]): if pd_client.capacity_share is not None else pd_client.start_args["running_max_req_size"] ) - pd_client.capacity_share = max(0, int(load_info.get("capacity_share", fallback_capacity))) + capacity_share = max(0, int(load_info.get("capacity_share", fallback_capacity))) + capacity_changed = fallback_capacity != capacity_share + pd_client.capacity_share = capacity_share pd_client.capacity_epoch = capacity_epoch + return pd_client.mode == "decode" and capacity_changed except Exception as e: logger.warning(f"udpate node load info failed, load_info: {load_info} error: {str(e)}") - return + return False def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index d11794c6a..5814b11f6 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -34,7 +34,7 @@ "lightllm_num_running_reqs": "Number of running requests", "lightllm_pd_master_admission_queue_size": "Number of requests waiting at the PD master admission queue", "lightllm_pd_master_admission_active_slots": "Number of decode slots leased by the PD master", - "lightllm_pd_master_admission_cold_capacity": "Current cold-request slot capacity at the PD master", + "lightllm_pd_master_admission_decode_capacity": "Decode slot capacity currently assigned to the PD master", } @@ -116,7 +116,7 @@ def init_metrics(self, args): self.create_gauge("lightllm_num_running_reqs") self.create_gauge("lightllm_pd_master_admission_queue_size") self.create_gauge("lightllm_pd_master_admission_active_slots") - self.create_gauge("lightllm_pd_master_admission_cold_capacity") + self.create_gauge("lightllm_pd_master_admission_decode_capacity") def create_histogram(self, name, buckets, labelnames=None): all_labels = ["model_name"] + (labelnames or []) diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index e36dc1e67..f60a470c6 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -11,6 +11,13 @@ logger = init_logger(__name__) +# Keep per-Master capacity metadata inside ``start_args`` so a new P/D node can +# still register with an older Master whose PD_Client_Obj rejects unknown +# top-level fields. New Masters extract these reserved protocol keys. +PD_MASTER_CAPACITY_SHARE_KEY = "__pd_master_capacity_share" +PD_MASTER_CAPACITY_EPOCH_KEY = "__pd_master_capacity_epoch" + + # 节点的行为 class NodeRole(enum.Enum): P = "prefill" @@ -47,10 +54,6 @@ class ObjType(enum.Enum): @dataclass class _PD_Client_RunStatus: total_token_usage_rate: float = 0.0 # pd 节点上的 token 使用率 - radix_cache_total_tokens: int = 0 - radix_cache_refed_tokens: int = 0 - radix_cache_capacity_tokens: int = 0 - report_time: float = 0.0 @dataclass diff --git a/unit_tests/server/test_pd_admission.py b/unit_tests/server/test_pd_admission.py index 80a623955..802e293c0 100644 --- a/unit_tests/server/test_pd_admission.py +++ b/unit_tests/server/test_pd_admission.py @@ -6,7 +6,6 @@ AdmissionPolicy, AdmissionPriority, AdmissionRequest, - CacheCapacitySnapshot, PDAdmissionController, SessionTracker, ) @@ -17,13 +16,11 @@ def _request( priority=AdmissionPriority.COLD, session_key=None, decode_slots=1, - estimated_uncached_work=0, ): return AdmissionRequest( session_key=session_key, priority=priority, decode_slots=decode_slots, - estimated_uncached_work=estimated_uncached_work, ) @@ -74,6 +71,45 @@ async def run(): asyncio.run(run()) +def test_decode_capacity_one_still_follows_slot_weighted_drr(): + async def run(): + policy = AdmissionPolicy( + continuation_weight=8, + probable_cache_hit_weight=3, + cold_weight=1, + waiting_decode_waves=12, + ) + controller = PDAdmissionController(lambda: 1, policy=policy) + active = await controller.acquire(_request()) + acquired = asyncio.Queue() + + async def acquire(label, priority): + lease = await controller.acquire(_request(priority)) + await acquired.put((label, lease)) + + tasks = [ + asyncio.create_task(acquire(f"continuation-{index}", AdmissionPriority.CONTINUATION)) for index in range(8) + ] + tasks.extend( + asyncio.create_task(acquire(f"probable-{index}", AdmissionPriority.PROBABLE_CACHE_HIT)) + for index in range(3) + ) + tasks.append(asyncio.create_task(acquire("cold-0", AdmissionPriority.COLD))) + await asyncio.sleep(0) + + active.release() + expected = ["continuation"] * 8 + ["probable"] * 3 + ["cold"] + for expected_prefix in expected: + label, lease = await asyncio.wait_for(acquired.get(), timeout=1.0) + assert label.startswith(expected_prefix) + lease.release() + + await asyncio.gather(*tasks) + assert controller.active_slots == 0 + + asyncio.run(run()) + + def test_higher_priority_request_can_replace_a_queued_cold_request(): async def run(): controller = PDAdmissionController(lambda: 1) @@ -111,7 +147,7 @@ async def run(): asyncio.run(run()) -def test_multi_choice_request_reserves_capacity_across_individual_releases(): +def test_same_priority_small_request_backfills_a_blocked_gang_without_idle_slots(): async def run(): controller = PDAdmissionController(lambda: 3) active = [await controller.acquire(_request()) for _ in range(3)] @@ -120,52 +156,138 @@ async def run(): await asyncio.sleep(0) active[0].release() - await asyncio.sleep(0) + later = await later_task assert multi_choice_task.done() is False - assert later_task.done() is False + assert controller.active_slots == 3 + assert all(deficit >= 0 for deficit in controller._deficits.values()) active[1].release() - multi_choice = await multi_choice_task - assert later_task.done() is False + assert multi_choice_task.done() is False - active[2].release() - later = await later_task - multi_choice.release() later.release() + multi_choice = await multi_choice_task + assert controller.active_slots == 3 + multi_choice.release() + active[2].release() asyncio.run(run()) -def test_multi_choice_reservation_does_not_block_fitting_higher_priority_request(): +def test_blocked_gang_keeps_backfill_open_for_a_later_small_request(): async def run(): - controller = PDAdmissionController(lambda: 3) - active = [await controller.acquire(_request()) for _ in range(3)] - - # 先消费调度表中的 continuation 和 probable 配额,使下一次轮到 cold。 - first_continuation_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) - first_probable_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.PROBABLE_CACHE_HIT))) + controller = PDAdmissionController( + lambda: 2, + policy=AdmissionPolicy(waiting_decode_waves=2), + ) + active = [await controller.acquire(_request()) for _ in range(2)] + gang_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) await asyncio.sleep(0) + active[0].release() - first_continuation = await first_continuation_task + assert gang_task.done() is False + assert controller.active_slots == 1 + + small_task = asyncio.create_task(controller.acquire(_request())) + small = await small_task + assert controller.active_slots == 2 + assert gang_task.done() is False + + small.release() + assert gang_task.done() is False active[1].release() - first_probable = await first_probable_task + gang = await gang_task + gang.release() + + asyncio.run(run()) - large_cold_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) - later_continuation_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + +def test_full_queue_of_session_blocked_waiters_cannot_leave_other_sessions_idle(): + async def run(): + controller = PDAdmissionController(lambda: 4) + active_a = await controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="session-a")) + queued_a_tasks = [ + asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="session-a"))) + for _ in range(4) + ] await asyncio.sleep(0) + assert controller.queued_slots == 4 + assert controller.active_slots == 1 - active[2].release() - later_continuation = await later_continuation_task + session_b = await controller.acquire(_request(AdmissionPriority.CONTINUATION, session_key="session-b")) + assert controller.active_slots == 2 + assert controller.queued_slots == 4 + + session_b.release() + active_a.release() + for task in queued_a_tasks: + lease = await task + lease.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_full_gang_queue_still_accepts_a_fitting_request_within_backfill_budget(): + async def run(): + controller = PDAdmissionController(lambda: 4) + active = await controller.acquire(_request(decode_slots=3)) + gang_tasks = [asyncio.create_task(controller.acquire(_request(decode_slots=2))) for _ in range(2)] + await asyncio.sleep(0) + assert controller.queued_slots == 4 assert controller.active_slots == 3 - assert large_cold_task.done() is False - first_continuation.release() - assert large_cold_task.done() is False - first_probable.release() - large_cold = await large_cold_task + small = await controller.acquire(_request()) + assert controller.active_slots == 4 + assert controller.queued_slots == 4 + assert controller._backfilled_slots == 1 - later_continuation.release() - large_cold.release() + small.release() + active.release() + first_gang = await gang_tasks[0] + second_gang = await gang_tasks[1] + first_gang.release() + second_gang.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_gang_reservation_starts_after_one_backfill_wave_and_blocks_later_priority(): + async def run(): + controller = PDAdmissionController( + lambda: 3, + policy=AdmissionPolicy(waiting_decode_waves=5), + ) + active = [await controller.acquire(_request()) for _ in range(3)] + gang_task = asyncio.create_task(controller.acquire(_request(decode_slots=3))) + backfill_tasks = [asyncio.create_task(controller.acquire(_request())) for _ in range(3)] + await asyncio.sleep(0) + + backfills = [] + for active_lease, backfill_task in zip(active, backfill_tasks): + active_lease.release() + backfills.append(await backfill_task) + assert gang_task.done() is False + + later_high_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + + backfills[0].release() + await asyncio.sleep(0) + assert later_high_task.done() is False + assert gang_task.done() is False + backfills[1].release() + await asyncio.sleep(0) + assert later_high_task.done() is False + assert gang_task.done() is False + + backfills[2].release() + gang = await gang_task + assert later_high_task.done() is False + + gang.release() + later_high = await later_high_task + later_high.release() asyncio.run(run()) @@ -210,6 +332,38 @@ async def run(): asyncio.run(run()) +def test_cancelling_reserved_gang_removes_the_backfill_barrier(): + async def run(): + controller = PDAdmissionController( + lambda: 2, + policy=AdmissionPolicy(waiting_decode_waves=4), + ) + active = [await controller.acquire(_request()) for _ in range(2)] + gang_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) + backfill_tasks = [asyncio.create_task(controller.acquire(_request())) for _ in range(2)] + await asyncio.sleep(0) + + backfills = [] + for active_lease, backfill_task in zip(active, backfill_tasks): + active_lease.release() + backfills.append(await backfill_task) + + later_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + assert later_task.done() is False + + gang_task.cancel() + with pytest.raises(asyncio.CancelledError): + await gang_task + + backfills[0].release() + later = await later_task + backfills[1].release() + later.release() + + asyncio.run(run()) + + def test_wait_timeout_removes_request_from_queue(): async def run(): policy = AdmissionPolicy( @@ -229,6 +383,40 @@ async def run(): asyncio.run(run()) +def test_reserved_gang_timeout_removes_the_backfill_barrier(): + async def run(): + policy = AdmissionPolicy( + continuation_max_wait_seconds=1.0, + probable_cache_hit_max_wait_seconds=1.0, + cold_max_wait_seconds=0.03, + waiting_decode_waves=4, + ) + controller = PDAdmissionController(lambda: 2, policy=policy) + active = [await controller.acquire(_request()) for _ in range(2)] + gang_task = asyncio.create_task(controller.acquire(_request(decode_slots=2))) + await asyncio.sleep(0) + + active[0].release() + first_backfill_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + first_backfill = await first_backfill_task + + second_backfill_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + await asyncio.sleep(0) + active[1].release() + second_backfill = await second_backfill_task + + later_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) + with pytest.raises(ServerBusyError, match="wait timed out"): + await gang_task + + first_backfill.release() + later = await later_task + second_backfill.release() + later.release() + + asyncio.run(run()) + + def test_capacity_changes_wake_waiters_without_overcommitting(): async def run(): capacity = [1] @@ -258,6 +446,99 @@ async def run(): asyncio.run(run()) +def test_capacity_shrink_fails_an_oversized_blocked_gang_and_clears_state(): + async def run(): + capacity = [3] + controller = PDAdmissionController( + lambda: capacity[0], + policy=AdmissionPolicy(waiting_decode_waves=2), + ) + active = [await controller.acquire(_request()) for _ in range(3)] + gang_task = asyncio.create_task(controller.acquire(_request(decode_slots=3))) + await asyncio.sleep(0) + + active[0].release() + assert gang_task.done() is False + + capacity[0] = 2 + controller.on_capacity_change() + with pytest.raises(ServerBusyError, match="fell below queued request size"): + await gang_task + assert controller.queued_slots == 0 + + small_task = asyncio.create_task(controller.acquire(_request())) + await asyncio.sleep(0) + assert small_task.done() is False + active[1].release() + small = await small_task + active[2].release() + small.release() + + asyncio.run(run()) + + +def test_capacity_shrink_trims_queue_to_dynamic_slot_limit_by_priority_and_recency(): + async def run(): + capacity = [4] + policy = AdmissionPolicy(waiting_decode_waves=2) + controller = PDAdmissionController(lambda: capacity[0], policy=policy) + active = await controller.acquire(_request(decode_slots=4)) + + cold_tasks = [asyncio.create_task(controller.acquire(_request(AdmissionPriority.COLD))) for _ in range(3)] + probable_tasks = [ + asyncio.create_task(controller.acquire(_request(AdmissionPriority.PROBABLE_CACHE_HIT))) for _ in range(2) + ] + continuation_tasks = [ + asyncio.create_task(controller.acquire(_request(AdmissionPriority.CONTINUATION))) for _ in range(3) + ] + await asyncio.sleep(0) + assert controller.queued_slots == 8 + + capacity[0] = 1 + controller.on_capacity_change() + assert controller.queued_slots == capacity[0] * policy.waiting_decode_waves + + for task in cold_tasks + probable_tasks + [continuation_tasks[-1]]: + with pytest.raises(ServerBusyError, match="queue capacity shrank"): + await task + assert continuation_tasks[0].done() is False + assert continuation_tasks[1].done() is False + + active.release() + first = await continuation_tasks[0] + assert continuation_tasks[1].done() is False + first.release() + second = await continuation_tasks[1] + second.release() + + asyncio.run(run()) + + +def test_zero_capacity_fails_all_queued_requests(): + async def run(): + capacity = [2] + controller = PDAdmissionController( + lambda: capacity[0], + policy=AdmissionPolicy(waiting_decode_waves=2), + ) + active = [await controller.acquire(_request()) for _ in range(2)] + tasks = [asyncio.create_task(controller.acquire(_request())) for _ in range(2)] + await asyncio.sleep(0) + + capacity[0] = 0 + controller.on_capacity_change() + for task in tasks: + with pytest.raises(ServerBusyError, match="fell below queued request size"): + await task + assert controller.queued_slots == 0 + assert controller.queued_request_count == 0 + + for lease in active: + lease.release() + + asyncio.run(run()) + + def test_session_tracker_requires_a_recent_observed_success(): now = [0.0] tracker = SessionTracker(ttl_seconds=10, max_sessions=2, clock=lambda: now[0]) @@ -325,103 +606,22 @@ async def run(): asyncio.run(run()) -def test_cold_capacity_tracks_live_cache_headroom_and_actual_miss_size(): - snapshot = [CacheCapacitySnapshot(total_tokens=200, capacity_tokens=1000)] - controller = PDAdmissionController( - lambda: 4, - cache_capacity_provider=lambda: snapshot[0], - ) - cold = _request(AdmissionPriority.COLD) - - controller.record_prefill_result(cold, prompt_tokens=100, cached_tokens=0) - assert controller.cold_capacity == 4 - - snapshot[0] = CacheCapacitySnapshot(total_tokens=800, capacity_tokens=1000) - controller.on_capacity_change() - assert controller.cold_capacity == 2 - - snapshot[0] = CacheCapacitySnapshot(total_tokens=1000, capacity_tokens=1000) - controller.on_capacity_change() - assert controller.cold_capacity == 1 - - -def test_cache_friendly_request_bypasses_cold_capacity(): - async def run(): - snapshot = CacheCapacitySnapshot(total_tokens=1000, capacity_tokens=1000) - controller = PDAdmissionController( - lambda: 3, - cache_capacity_provider=lambda: snapshot, - ) - cold_request = _request(AdmissionPriority.COLD) - controller.record_prefill_result(cold_request, prompt_tokens=100, cached_tokens=0) - - first_cold = await controller.acquire(cold_request) - second_cold_task = asyncio.create_task(controller.acquire(cold_request)) - await asyncio.sleep(0) - assert second_cold_task.done() is False - - probable = await controller.acquire(_request(AdmissionPriority.PROBABLE_CACHE_HIT)) - assert controller.active_slots == 2 - probable.release() - - first_cold.release() - second_cold = await second_cold_task - second_cold.release() - - asyncio.run(run()) - - -def test_probable_cache_hits_consume_cold_capacity_when_actual_hits_are_low(): +def test_requests_remain_fifo_within_the_same_priority_class(): async def run(): - snapshot = CacheCapacitySnapshot(total_tokens=1000, capacity_tokens=1000) controller = PDAdmissionController( - lambda: 3, - cache_capacity_provider=lambda: snapshot, - ) - probable_request = _request(AdmissionPriority.PROBABLE_CACHE_HIT) - - initially_trusted = await controller.acquire(probable_request) - assert controller.active_cold_slots == 0 - controller.record_prefill_result( - probable_request, - prompt_tokens=100, - cached_tokens=0, + lambda: 1, + policy=AdmissionPolicy(waiting_decode_waves=2), ) - assert controller.cold_capacity == 1 - - first_gated = await controller.acquire(probable_request) - assert controller.active_cold_slots == 1 - second_gated_task = asyncio.create_task(controller.acquire(probable_request)) - await asyncio.sleep(0) - assert second_gated_task.done() is False - - # 可信度变化不能让已经取得的租约在释放时误扣冷槽位。 - initially_trusted.release() - assert controller.active_cold_slots == 1 - assert second_gated_task.done() is False - - first_gated.release() - second_gated = await second_gated_task - second_gated.release() - - asyncio.run(run()) - - -def test_smaller_cold_request_is_dispatched_first(): - async def run(): - controller = PDAdmissionController(lambda: 2) - active = [await controller.acquire(_request()) for _ in range(2)] - large_task = asyncio.create_task(controller.acquire(_request(estimated_uncached_work=100))) - small_task = asyncio.create_task(controller.acquire(_request(estimated_uncached_work=10))) + active = await controller.acquire(_request()) + first_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.COLD, session_key="first"))) + second_task = asyncio.create_task(controller.acquire(_request(AdmissionPriority.COLD, session_key="second"))) await asyncio.sleep(0) - active[0].release() - small = await small_task - assert large_task.done() is False - - active[1].release() - large = await large_task - small.release() - large.release() + active.release() + first = await first_task + assert second_task.done() is False + first.release() + second = await second_task + second.release() asyncio.run(run()) diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index f3467fa70..252d7b0c7 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -1,12 +1,17 @@ import asyncio import json from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, call import pytest from easydict import EasyDict from lightllm.server.core.objs.start_args_type import StartArgs -from lightllm.server.httpserver.pd_loop import _allocate_capacity_share, _update_pd_master_membership +from lightllm.server.httpserver.pd_loop import ( + _allocate_capacity_share, + _build_pd_registration_info, + _update_pd_master_membership, +) from lightllm.server.httpserver_for_pd_master.admission import ( AdmissionPolicy, AdmissionPriority, @@ -14,6 +19,7 @@ SessionTracker, ) from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager +from lightllm.server.pd_io_struct import PD_MASTER_CAPACITY_EPOCH_KEY, PD_MASTER_CAPACITY_SHARE_KEY def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): @@ -200,6 +206,52 @@ def test_pd_node_capacity_is_partitioned_without_overlap(): assert _allocate_capacity_share(8, master_ids, 99) == 0 +def test_pd_registration_keeps_legacy_top_level_schema(monkeypatch): + from lightllm.server.httpserver import pd_loop + + args = SimpleNamespace(pd_node_id=7, running_max_req_size=8, host="0.0.0.0") + manager = SimpleNamespace( + args=args, + host_ip="10.0.0.7", + pd_mode=SimpleNamespace(value="decode"), + pd_master_ids=(10, 20), + pd_master_capacity_epoch=123, + ) + monkeypatch.setattr(pd_loop, "get_shm_port_args", lambda: SimpleNamespace(port=8001)) + + registration = _build_pd_registration_info(manager, SimpleNamespace(node_id=10)) + + assert set(registration) == {"node_id", "client_ip_port", "mode", "start_args"} + assert registration["start_args"][PD_MASTER_CAPACITY_SHARE_KEY] == 4 + assert registration["start_args"][PD_MASTER_CAPACITY_EPOCH_KEY] == 123 + assert args.host == "0.0.0.0" + + +def test_pd_node_load_info_omits_radix_cache_hot_path(monkeypatch): + from lightllm.server import api_http + from lightllm.server.httpserver import pd_loop + + args = SimpleNamespace(tp=4, dp=2, nnodes=1, running_max_req_size=8) + httpserver_manager = SimpleNamespace( + host_ip="10.0.0.1", + pd_master_ids=(10, 20), + pd_master_capacity_epoch=123, + ) + shared_token_load = SimpleNamespace(get_dynamic_max_load=lambda dp_index: (0.25, 0.75)[dp_index]) + monkeypatch.setattr(api_http.g_objs, "args", args) + monkeypatch.setattr(api_http.g_objs, "httpserver_manager", httpserver_manager) + monkeypatch.setattr(api_http.g_objs, "shared_token_load", shared_token_load) + monkeypatch.setattr(pd_loop, "get_shm_port_args", lambda: SimpleNamespace(port=8001)) + + assert not hasattr(pd_loop, "_get_radix_cache_info") + assert pd_loop._get_load_info(pd_master_node_id=10) == { + "total_token_usage_rate": 0.5, + "client_ip_port": "10.0.0.1:8001", + "capacity_share": 4, + "capacity_epoch": 123, + } + + def test_pd_master_membership_change_advances_epoch_and_wakes_heartbeats(monkeypatch): async def run(): manager = SimpleNamespace() @@ -222,7 +274,7 @@ async def run(): asyncio.run(run()) -def test_pd_manager_uses_latest_decode_capacity_lease_and_cache_telemetry(): +def test_pd_manager_reports_only_actual_decode_capacity_changes(): args = StartArgs() manager = PDManager(args) client_ip_port = "10.0.0.2:8000" @@ -243,36 +295,133 @@ def test_pd_manager_uses_latest_decode_capacity_lease_and_cache_telemetry(): assert manager.get_decode_capacity() == 3 - manager.update_node_load_info( - { - "client_ip_port": client_ip_port, - "total_token_usage_rate": 0.25, - "capacity_share": 1, - "capacity_epoch": 99, - } + assert ( + manager.update_node_load_info( + { + "client_ip_port": client_ip_port, + "total_token_usage_rate": 0.25, + "capacity_share": 1, + "capacity_epoch": 99, + } + ) + is False ) assert manager.get_decode_capacity() == 3 - manager.update_node_load_info( - { - "client_ip_port": client_ip_port, - "total_token_usage_rate": 0.5, - "capacity_share": 2, - "capacity_epoch": 101, - "radix_cache_total_tokens": 700, - "radix_cache_refed_tokens": 200, - "radix_cache_capacity_tokens": 1000, - } + assert ( + manager.update_node_load_info( + { + "client_ip_port": client_ip_port, + "total_token_usage_rate": 0.5, + "capacity_share": 2, + "capacity_epoch": 101, + # Old nodes may continue sending cache telemetry during a + # rolling upgrade. It must not affect decode admission. + "radix_cache_total_tokens": 700, + "radix_cache_refed_tokens": 200, + "radix_cache_capacity_tokens": 1000, + } + ) + is True ) node = manager.decode_nodes[0] assert manager.get_decode_capacity() == 2 assert node.run_status.total_token_usage_rate == 0.5 - assert node.run_status.radix_cache_total_tokens == 700 - assert node.run_status.radix_cache_refed_tokens == 200 - assert node.run_status.radix_cache_capacity_tokens == 1000 + + # A fresh report and a load-only change are not capacity changes. + assert ( + manager.update_node_load_info( + { + "client_ip_port": client_ip_port, + "total_token_usage_rate": 0.75, + "capacity_share": 2, + "capacity_epoch": 102, + } + ) + is False + ) + assert manager.get_decode_capacity() == 2 + + +def test_pd_registration_reads_capacity_from_legacy_safe_start_args(): + args = StartArgs() + manager = PDManager(args) + manager.register_pd( + { + "node_id": 2, + "client_ip_port": "10.0.0.2:8000", + "mode": "decode", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + PD_MASTER_CAPACITY_SHARE_KEY: 3, + PD_MASTER_CAPACITY_EPOCH_KEY: 100, + }, + }, + websocket=object(), + ) + + node = manager.decode_nodes[0] + assert node.capacity_share == 3 + assert node.capacity_epoch == 100 + assert manager.get_decode_capacity() == 3 -def test_prefill_cache_headroom_is_scaled_to_the_master_decode_lease(): +def test_pd_disconnect_accepts_transitional_top_level_capacity_fields(): + args = StartArgs() + manager = PDManager(args) + registration = { + "node_id": 2, + "client_ip_port": "10.0.0.2:8000", + "mode": "decode", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + }, + "capacity_share": 3, + "capacity_epoch": 100, + } + manager.register_pd(registration, websocket=object()) + + manager.remove_pd(registration) + + assert manager.decode_nodes == [] + assert manager.url_to_pd_nodes == {} + + +def test_materializing_decode_fallback_share_is_not_a_capacity_change(): + args = StartArgs() + manager = PDManager(args) + client_ip_port = "10.0.0.2:8000" + manager.register_pd( + { + "node_id": 2, + "client_ip_port": client_ip_port, + "mode": "decode", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "running_max_req_size": 8, + }, + }, + websocket=object(), + ) + + assert manager.decode_nodes[0].capacity_share is None + assert manager.get_decode_capacity() == 8 + assert ( + manager.update_node_load_info( + { + "client_ip_port": client_ip_port, + "total_token_usage_rate": 0.5, + "capacity_epoch": 1, + } + ) + is False + ) + assert manager.get_decode_capacity() == 8 + + +def test_prefill_cache_telemetry_does_not_change_decode_capacity(): args = StartArgs() manager = PDManager(args) manager.register_pd( @@ -302,21 +451,201 @@ def test_prefill_cache_headroom_is_scaled_to_the_master_decode_lease(): }, websocket=object(), ) - manager.update_node_load_info( - { - "client_ip_port": "10.0.0.1:8000", - "total_token_usage_rate": 0.25, - "radix_cache_total_tokens": 800, - "radix_cache_refed_tokens": 100, - "radix_cache_capacity_tokens": 1000, - } + + assert manager.get_decode_capacity() == 4 + assert ( + manager.update_node_load_info( + { + "client_ip_port": "10.0.0.1:8000", + "total_token_usage_rate": 0.25, + "radix_cache_total_tokens": 1000, + "radix_cache_refed_tokens": 100, + "radix_cache_capacity_tokens": 1000, + } + ) + is False ) + assert manager.get_decode_capacity() == 4 + - snapshot = manager.get_prefill_cache_capacity() - assert snapshot is not None - assert snapshot.total_tokens == 350 - assert snapshot.capacity_tokens == 375 - assert snapshot.free_tokens == 25 +def test_pd_master_redrains_admission_only_for_decode_capacity_changes(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.pd_manager = SimpleNamespace(update_node_load_info=MagicMock(side_effect=[False, True])) + manager.admission_controller = SimpleNamespace(on_capacity_change=MagicMock()) + + unchanged_load = {"client_ip_port": "10.0.0.1:8000", "capacity_share": 4} + changed_load = {"client_ip_port": "10.0.0.1:8000", "capacity_share": 2} + + manager.update_node_load_info(unchanged_load) + manager.admission_controller.on_capacity_change.assert_not_called() + + manager.update_node_load_info(changed_load) + manager.admission_controller.on_capacity_change.assert_called_once_with() + assert manager.pd_manager.update_node_load_info.call_args_list == [ + call(unchanged_load), + call(changed_load), + ] + + +def test_token_packs_redrain_admission_only_after_decode_capacity_change(): + from lightllm.server.pd_io_struct import ObjType + + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = SimpleNamespace(config_server_host=None) + manager.pd_manager = SimpleNamespace(update_node_load_info=MagicMock(side_effect=[False, True])) + manager.admission_controller = SimpleNamespace(on_capacity_change=MagicMock()) + manager.timer_log = AsyncMock() + manager.infos_queues = None + manager.req_id_to_out_inf = {} + + handle_task = asyncio.create_task(manager.handle_loop()) + try: + while manager.infos_queues is None: + await asyncio.sleep(0) + + unchanged_load = {"client_ip_port": "10.0.0.1:8000", "capacity_share": 4} + changed_load = {"client_ip_port": "10.0.0.1:8000", "capacity_share": 2} + await manager.put_to_handle_queue((ObjType.TOKEN_PACKS, [], unchanged_load)) + await manager.put_to_handle_queue((ObjType.TOKEN_PACKS, [], changed_load)) + + for _ in range(100): + if manager.pd_manager.update_node_load_info.call_count == 2: + break + await asyncio.sleep(0) + + assert manager.pd_manager.update_node_load_info.call_args_list == [ + call(unchanged_load), + call(changed_load), + ] + manager.admission_controller.on_capacity_change.assert_called_once_with() + finally: + handle_task.cancel() + await asyncio.gather(handle_task, return_exceptions=True) + + asyncio.run(run()) + + +def test_pd_master_deduplicates_unchanged_admission_metrics(monkeypatch): + from lightllm.server.httpserver_for_pd_master import manager as manager_module + + metric_client = SimpleNamespace(gauge_set=MagicMock()) + monkeypatch.setattr(manager_module, "MetricClient", lambda _port: metric_client) + monkeypatch.setattr(manager_module, "ReqIDGenerator", lambda: object()) + monkeypatch.setattr(manager_module, "get_tokenizer", lambda *_args, **_kwargs: object()) + monkeypatch.setattr(manager_module, "get_shm_port_args", lambda: SimpleNamespace(metric_port=1234)) + + manager = HttpServerManagerForPDMaster(StartArgs(max_req_total_len=1024)) + manager.pd_manager.get_decode_capacity = lambda: 4 + controller = SimpleNamespace(queued_request_count=2, active_slots=1) + + async def run(): + manager._record_admission_state(controller) + controller.active_slots = 2 + manager._record_admission_state(controller) + assert metric_client.gauge_set.call_count == 0 + + await asyncio.sleep(0) + assert {metric_call.args for metric_call in metric_client.gauge_set.call_args_list} == { + ("lightllm_pd_master_admission_queue_size", 2), + ("lightllm_pd_master_admission_active_slots", 2), + ("lightllm_pd_master_admission_decode_capacity", 4), + } + first_call_count = metric_client.gauge_set.call_count + + manager._record_admission_state(controller) + await asyncio.sleep(0) + assert metric_client.gauge_set.call_count == first_call_count + + controller.active_slots = 3 + manager._record_admission_state(controller) + await asyncio.sleep(0) + assert metric_client.gauge_set.call_args_list[-1] == call("lightllm_pd_master_admission_active_slots", 3) + assert metric_client.gauge_set.call_count == first_call_count + 1 + + asyncio.run(run()) + + +def test_pd_master_decode_lease_covers_prefill_and_stream_lifecycle(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs() + manager.pd_manager = SimpleNamespace( + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: 0.0), + ) + manager.admission_policy = AdmissionPolicy() + manager.admission_controller = PDAdmissionController(lambda: 1, policy=manager.admission_policy) + manager.session_tracker = SessionTracker( + ttl_seconds=manager.admission_policy.active_session_ttl_seconds, + max_sessions=manager.admission_policy.max_tracked_sessions, + ) + manager.metric_client = SimpleNamespace(histogram_observe=lambda *_args: None) + manager.running_request_count = 0 + manager.latest_success_infer_time = 0 + dispatched_prompts = [] + + async def fake_generate(prompt, *_args): + dispatched_prompts.append(prompt) + yield 1, "prefill", {"prompt_tokens": 16, "prompt_cache_len": 0}, object() + yield 1, "decode", {}, object() + + manager._generate = fake_generate + first = manager.generate("first", None, None, None) + first_prefill = await first.__anext__() + assert first_prefill[1] == "prefill" + assert manager.admission_controller.active_slots == 1 + + second = manager.generate("second", None, None, None) + second_result = asyncio.create_task(second.__anext__()) + await asyncio.sleep(0) + assert second_result.done() is False + assert dispatched_prompts == ["first"] + + first_decode = await first.__anext__() + assert first_decode[1] == "decode" + assert manager.admission_controller.active_slots == 1 + await first.aclose() + + assert (await second_result)[1] == "prefill" + await second.aclose() + assert manager.admission_controller.active_slots == 0 + + asyncio.run(run()) + + +def test_pd_master_releases_admission_lease_when_post_acquire_setup_fails(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = StartArgs() + manager.pd_manager = SimpleNamespace( + selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: 0.0), + ) + manager.admission_policy = AdmissionPolicy() + manager.admission_controller = PDAdmissionController(lambda: 1, policy=manager.admission_policy) + manager.session_tracker = SessionTracker( + ttl_seconds=manager.admission_policy.active_session_ttl_seconds, + max_sessions=manager.admission_policy.max_tracked_sessions, + ) + + def fail_histogram(*_args): + raise RuntimeError("metric enqueue failed") + + manager.metric_client = SimpleNamespace(histogram_observe=fail_histogram) + manager.running_request_count = 0 + manager.latest_success_infer_time = 0 + + async def fake_generate(*_args): + yield "must not dispatch" + + manager._generate = fake_generate + generator = manager.generate("prompt", SimpleNamespace(n=1), None, None) + with pytest.raises(RuntimeError, match="metric enqueue failed"): + await generator.__anext__() + + assert manager.admission_controller.active_slots == 0 + assert manager.running_request_count == 0 + + asyncio.run(run()) def test_prefill_registration_preserves_existing_inflight_prompt_chars(): @@ -438,7 +767,7 @@ async def fake_generate(prompt, *_args): asyncio.run(run()) -def test_pd_master_admission_classifies_session_cache_and_multi_choice_cost(): +def test_pd_master_admission_classifies_priority_and_multi_choice_cost(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.pd_manager = SimpleNamespace( selector=SimpleNamespace(estimate_prompt_cache_hit_rate=lambda _prompt: 0.75), @@ -454,7 +783,6 @@ def test_pd_master_admission_classifies_session_cache_and_multi_choice_cost(): ) assert probable.priority == AdmissionPriority.PROBABLE_CACHE_HIT assert probable.decode_slots == 3 - assert probable.estimated_uncached_work == 9 manager.session_tracker.mark_success("session-a") continuation = manager._build_admission_request(