From 65bde0abc0d28a59e22622000d77b9fe1a19ed4a Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 31 Aug 2026 13:50:41 +0800 Subject: [PATCH 01/10] 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 60b5fad4e8..447fa2e5e4 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 a9aef608bd..0f14e53d84 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 b422bf7703..4ac9ffc3f1 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 adf911bc21..08d7a6e904 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 5474806b7b..a8c3a6e046 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 a3dc268a93..da07759a99 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 6bdde78576..45fe091731 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 edb758a94b..fed4563f0d 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 02/10] 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 08d7a6e904..b143d41beb 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 45fe091731..1e141110d5 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 03/10] 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 447fa2e5e4..60b5fad4e8 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 0f14e53d84..a9aef608bd 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 4ac9ffc3f1..58dba3423c 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 fed4563f0d..b7e06b140a 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 04/10] 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 58dba3423c..6931dd0d3a 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 b143d41beb..4a5357c937 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 a8c3a6e046..7bf8443415 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 05/10] 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 60b5fad4e8..16739babff 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 1d8b2112fc..26c27e6fe0 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 e66540cb5e..61656459af 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 0000000000..6cda68274e --- /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 6931dd0d3a..ea26bd328b 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 0d42462c3f..d11794c6a9 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 8a1e2bd42b..e36dc1e677 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 0000000000..92d0760a07 --- /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 b7e06b140a..f3467fa701 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 06/10] 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 ea26bd328b..0624b44bea 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 07/10] 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 6cda68274e..94ea62e539 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 0624b44bea..6bf32108bb 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 92d0760a07..80a623955b 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 08/10] 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 61656459af..c3e3110b47 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 94ea62e539..9f25ebed83 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 6bf32108bb..675ba35435 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 4a5357c937..6bb9bdf6ef 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 7bf8443415..eb0f9b5704 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 09/10] 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 16739babff..6500ad53b1 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 c3e3110b47..95a280bdf4 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 9f25ebed83..bebd026869 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 675ba35435..fba20821d4 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 d11794c6a9..5814b11f61 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 e36dc1e677..f60a470c6e 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 80a623955b..802e293c0f 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 f3467fa701..252d7b0c75 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( From ff7d5a691ddeab376121f67fde917aa9e849c7f6 Mon Sep 17 00:00:00 2001 From: sufubao Date: Wed, 2 Sep 2026 00:31:27 +0800 Subject: [PATCH 10/10] refactor(pd): move decode admission to nodes --- lightllm/server/api_cli.py | 25 +- lightllm/server/api_http_pd.py | 2 - lightllm/server/api_start.py | 5 + lightllm/server/core/objs/shm_req_manager.py | 25 +- lightllm/server/core/objs/start_args_type.py | 3 + .../server/httpserver/decode_admission.py | 180 +++++ lightllm/server/httpserver/manager.py | 105 ++- lightllm/server/httpserver/pd_loop.py | 217 ++++-- .../httpserver_for_pd_master/admission.py | 716 ------------------ .../httpserver_for_pd_master/manager.py | 378 +++++---- .../pd_selector/cache_aware.py | 44 +- .../pd_selector/pd_selector.py | 12 +- lightllm/server/metrics/metrics.py | 6 - lightllm/server/pd_io_struct.py | 18 +- .../test_pd_master_multi_choice.py | 144 +++- .../server/core/objs/test_shm_req_manager.py | 20 + .../httpserver/test_decode_admission.py | 169 +++++ .../httpserver/test_pd_generate_error.py | 157 +++- .../test_pd_master_cached_tokens.py | 3 +- .../test_running_request_lifecycle.py | 67 +- unit_tests/server/test_pd_admission.py | 627 --------------- unit_tests/server/test_pd_cache_aware.py | 65 -- unit_tests/server/test_pd_master_mode.py | 523 +------------ 23 files changed, 1201 insertions(+), 2310 deletions(-) create mode 100644 lightllm/server/httpserver/decode_admission.py delete mode 100644 lightllm/server/httpserver_for_pd_master/admission.py create mode 100644 unit_tests/server/httpserver/test_decode_admission.py delete mode 100644 unit_tests/server/test_pd_admission.py diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 6500ad53b1..72f2b8c891 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -71,7 +71,30 @@ 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 admission queue based on registered decode capacity.", + help=( + "Deprecated alias for --disable_pd_node_decode_admission. Set it consistently on PD Masters and " + "Decode nodes when intentionally allowing Decode nodes without authoritative admission." + ), + ) + parser.add_argument( + "--disable_pd_node_decode_admission", + action="store_true", + help="Disable the bounded request-slot admission queue on PD Decode nodes.", + ) + parser.add_argument( + "--pd_node_decode_admission_queue_size", + type=int, + default=None, + help=( + "Maximum number of Decode request slots allowed to wait on a PD Decode node. " + "Defaults to --running_max_req_size." + ), + ) + parser.add_argument( + "--pd_node_decode_admission_timeout", + type=float, + default=5.0, + help="Maximum seconds a request may wait in the PD Decode node admission queue. Default: 5.", ) parser.add_argument( "--pd_trans_mode", diff --git a/lightllm/server/api_http_pd.py b/lightllm/server/api_http_pd.py index 26c27e6fe0..1d8b2112fc 100644 --- a/lightllm/server/api_http_pd.py +++ b/lightllm/server/api_http_pd.py @@ -41,8 +41,6 @@ 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/api_start.py b/lightllm/server/api_start.py index 37fe837ad1..8eac4f2dad 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -121,6 +121,11 @@ def _launch_subprocesses(args: StartArgs): args.mem_fraction > 0 and args.mem_fraction < 1 ), f"Invalid mem_fraction {args.mem_fraction}, The expected value is between 0 and 1." + if args.pd_node_decode_admission_queue_size is not None and args.pd_node_decode_admission_queue_size < 0: + raise ValueError("pd_node_decode_admission_queue_size must be non-negative") + if args.pd_node_decode_admission_timeout <= 0: + raise ValueError("pd_node_decode_admission_timeout must be positive") + if args.graph_max_len_in_batch == 0: args.graph_max_len_in_batch = args.max_req_total_len diff --git a/lightllm/server/core/objs/shm_req_manager.py b/lightllm/server/core/objs/shm_req_manager.py index c61376c8d0..4215e63d69 100644 --- a/lightllm/server/core/objs/shm_req_manager.py +++ b/lightllm/server/core/objs/shm_req_manager.py @@ -82,17 +82,34 @@ def init_alloc_state_shm(self): # alloc_req_index 和 release_req_index 是分配资源时使用的接口。 # 只有管理请求申请和释放的首节点才能调用这个接口。 def alloc_req_index(self): + indexes = self.alloc_req_indexes(1) + return None if indexes is None else indexes[0] + + def alloc_req_indexes(self, req_num: int): + """Atomically allocate all requested indexes or leave the pool unchanged.""" + if req_num < 1: + raise ValueError("req_num must be positive") + with self.manager_lock: - idx = self.linked_req_manager.alloc() - if idx is not None: + indexes = [] + for _ in range(req_num): + idx = self.linked_req_manager.alloc() + if idx is None: + for allocated_idx in reversed(indexes): + self.alloc_state_shm.arr[allocated_idx] = 0 + self.linked_req_manager.free(allocated_idx) + return None assert self.alloc_state_shm.arr[idx] == 0 self.alloc_state_shm.arr[idx] = 1 - return idx - return None + indexes.append(idx) + return indexes async def async_alloc_req_index(self): return self.alloc_req_index() + async def async_alloc_req_indexes(self, req_num: int): + return self.alloc_req_indexes(req_num) + def release_req_index(self, req_index_in_mem): assert req_index_in_mem < self.max_req_num with self.manager_lock: diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index a9aef608bd..76bb4715c9 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -23,6 +23,9 @@ 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) + disable_pd_node_decode_admission: bool = field(default=False) + pd_node_decode_admission_queue_size: Optional[int] = field(default=None) + pd_node_decode_admission_timeout: float = field(default=5.0) 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/decode_admission.py b/lightllm/server/httpserver/decode_admission.py new file mode 100644 index 0000000000..18fc3f1add --- /dev/null +++ b/lightllm/server/httpserver/decode_admission.py @@ -0,0 +1,180 @@ +from __future__ import annotations + +import asyncio +import time +from collections import deque +from dataclasses import dataclass +from typing import Callable, Deque, Optional + +from lightllm.utils.error_utils import ServerBusyError + + +@dataclass(slots=True) +class _Waiter: + slots: int + enqueue_time: float + future: asyncio.Future + + +class DecodeAdmissionLease: + """A fixed number of request slots reserved on one Decode node.""" + + def __init__(self, controller: "DecodeAdmissionController", slots: int, waited_seconds: float) -> None: + self._controller = controller + self.slots = slots + self.waited_seconds = waited_seconds + self._released = False + + def release(self) -> None: + if self._released: + return + self._released = True + self._controller._release(self.slots) + + def split(self, slots: list[int]) -> list["DecodeAdmissionLease"]: + """Transfer this reservation into independently releasable child leases.""" + if self._released: + raise RuntimeError("Decode admission lease has already been released") + if not slots or any(slot_count < 1 for slot_count in slots) or sum(slots) != self.slots: + raise ValueError("child lease slots must be positive and sum to the parent lease") + + self._released = True + return [DecodeAdmissionLease(self._controller, slot_count, self.waited_seconds) for slot_count in slots] + + +class DecodeAdmissionLeaseHandle: + """Transfers a pre-acquired lease into an asynchronously started generator.""" + + def __init__(self, lease: DecodeAdmissionLease) -> None: + self._lease: Optional[DecodeAdmissionLease] = lease + + def take(self) -> DecodeAdmissionLease: + if self._lease is None: + raise RuntimeError("Decode admission lease has already been taken") + lease = self._lease + self._lease = None + return lease + + def release(self) -> None: + if self._lease is None: + return + self._lease.release() + self._lease = None + + +class DecodeAdmissionController: + """A bounded, cancellable FIFO for request slots owned by one Decode node.""" + + def __init__( + self, + capacity: int, + max_queued_slots: int, + timeout_seconds: float, + clock: Callable[[], float] = time.monotonic, + ) -> None: + if capacity < 1: + raise ValueError("decode admission capacity must be positive") + if max_queued_slots < 0: + raise ValueError("decode admission queue size must be non-negative") + if timeout_seconds <= 0: + raise ValueError("decode admission timeout must be positive") + + self.capacity = capacity + self.max_queued_slots = max_queued_slots + self.timeout_seconds = timeout_seconds + self._clock = clock + self._active_slots = 0 + self._queued_slots = 0 + self._waiters: Deque[_Waiter] = deque() + + @property + def active_slots(self) -> int: + return self._active_slots + + @property + def queued_slots(self) -> int: + return self._queued_slots + + @property + def queued_request_count(self) -> int: + return len(self._waiters) + + async def acquire(self, slots: int) -> DecodeAdmissionLease: + if slots < 1: + raise ValueError("decode admission slots must be positive") + if slots > self.capacity: + raise ServerBusyError(f"request needs {slots} Decode slots, but the node capacity is {self.capacity}") + + if not self._waiters and self._active_slots + slots <= self.capacity: + return self._activate(slots, waited_seconds=0.0) + + if self._queued_slots + slots > self.max_queued_slots: + raise ServerBusyError("Decode node admission queue is full") + + waiter = _Waiter( + slots=slots, + enqueue_time=self._clock(), + future=asyncio.get_running_loop().create_future(), + ) + self._waiters.append(waiter) + self._queued_slots += slots + self._drain() + + try: + return await asyncio.wait_for(asyncio.shield(waiter.future), timeout=self.timeout_seconds) + except asyncio.TimeoutError as exc: + lease = self._remove_waiter_or_take_lease(waiter) + if lease is not None: + lease.release() + raise ServerBusyError("Decode node admission queue wait timed out") from exc + except BaseException: + lease = self._remove_waiter_or_take_lease(waiter) + if lease is not None: + lease.release() + raise + + def _activate(self, slots: int, waited_seconds: float) -> DecodeAdmissionLease: + self._active_slots += slots + return DecodeAdmissionLease(self, slots, waited_seconds) + + def _release(self, slots: int) -> None: + self._active_slots -= slots + if self._active_slots < 0: + raise RuntimeError("Decode admission active slot count became negative") + self._drain() + + def _drain(self) -> None: + while self._waiters: + waiter = self._waiters[0] + if waiter.future.done(): + self._remove_waiter(waiter) + continue + if self._active_slots + waiter.slots > self.capacity: + break + + self._remove_waiter(waiter) + waiter.future.set_result( + self._activate( + waiter.slots, + waited_seconds=max(0.0, self._clock() - waiter.enqueue_time), + ) + ) + + def _remove_waiter(self, waiter: _Waiter) -> bool: + try: + self._waiters.remove(waiter) + except ValueError: + return False + self._queued_slots -= waiter.slots + return True + + def _remove_waiter_or_take_lease(self, waiter: _Waiter) -> Optional[DecodeAdmissionLease]: + if self._remove_waiter(waiter): + waiter.future.cancel() + self._drain() + return None + if waiter.future.done() and not waiter.future.cancelled(): + result = waiter.future.result() + if isinstance(result, DecodeAdmissionLease): + return result + return None diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 9e6f77e58d..b00feb0a00 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -34,6 +34,7 @@ from lightllm.server.metrics.manager import MetricClient from .rl_controller import HttpRlController from .manager_ext import HttpRlManagerHelper +from .decode_admission import DecodeAdmissionController, DecodeAdmissionLease, DecodeAdmissionLeaseHandle from lightllm.utils.statics_utils import MovingAverage from lightllm.utils.config_utils import get_vocab_size from lightllm.utils.envs_utils import get_unique_server_name @@ -117,6 +118,19 @@ def __init__( self.pd_mode: NodeRole = NodeRole(self.args.run_mode) assert self.pd_mode in [NodeRole.NORMAL, NodeRole.P, NodeRole.D] + self.decode_admission_controller: Optional[DecodeAdmissionController] = None + decode_admission_disabled = ( + args.disable_pd_node_decode_admission or args.disable_pd_master_decode_capacity_limit + ) + if self.pd_mode.is_D() and not self.is_multinode_tp_slave and not decode_admission_disabled: + max_queued_slots = args.pd_node_decode_admission_queue_size + if max_queued_slots is None: + max_queued_slots = args.running_max_req_size + self.decode_admission_controller = DecodeAdmissionController( + capacity=args.running_max_req_size, + max_queued_slots=max_queued_slots, + timeout_seconds=args.pd_node_decode_admission_timeout, + ) self.id_gen = ReqIDGenerator() self.first_time_costs = MovingAverage() self.per_token_costs = MovingAverage() @@ -322,6 +336,9 @@ async def generate( pd_upload_websocket: ClientConnection = None, # 用于等待 pd_master 下发的交换信息 pd_event: asyncio.Event = None, + # Decode Master 可在所有 choice 完成 Prefill 后原子预留 n 个槽, + # 再通过该 handle 把每个 choice 的 lease 转交给本地请求生命周期。 + decode_admission_lease_handle: Optional[DecodeAdmissionLeaseHandle] = None, ) -> AsyncGenerator[Tuple[int, str, dict, FinishStatus], None]: start_time = time.time() @@ -338,6 +355,11 @@ async def generate( ) running_request_registered = False + decode_admission_lease: Optional[DecodeAdmissionLease] = None + if decode_admission_lease_handle is not None: + decode_admission_lease = decode_admission_lease_handle.take() + unmanaged_req_indexes = [] + unmanaged_req_objs = [] if not self.pd_mode.is_P(): await self._register_running_request() running_request_registered = True @@ -426,21 +448,26 @@ async def generate( await self._register_running_request() running_request_registered = True - # 申请资源并存储 - alloced_req_indexes = [] - while len(alloced_req_indexes) < sampling_params.n: - alloc_req_index = await self.shm_req_manager.async_alloc_req_index() - sleep_time = 0.1 - while alloc_req_index is None: - await asyncio.sleep(sleep_time) - sleep_time *= 1.1 - sleep_time = min(1, sleep_time) - - alloc_req_index = await self.shm_req_manager.async_alloc_req_index() - alloced_req_indexes.append(alloc_req_index) + # Decode admission happens immediately before the node reserves its + # local shm_req objects. All Masters therefore compete for the same + # authoritative pool, and an n-choice request reserves all n slots at once. + admission_controller = getattr(self, "decode_admission_controller", None) + if admission_controller is not None and decode_admission_lease is None: + decode_admission_lease = await admission_controller.acquire(sampling_params.n) + self.metric_client.histogram_observe( + "lightllm_request_queue_duration_bucket", + decode_admission_lease.waited_seconds, + ) + + alloced_req_indexes = await self._alloc_req_indexes( + sampling_params.n, + expect_available=decode_admission_lease is not None, + ) + unmanaged_req_indexes.extend(alloced_req_indexes) req_objs: List[Req] = [] for i, req_index in enumerate(alloced_req_indexes): req_obj = await self.shm_req_manager.async_get_req_obj_by_index(req_index) + unmanaged_req_objs.append(req_obj) req_obj.init( group_request_id + i, prompt_ids, @@ -461,8 +488,17 @@ async def generate( f"{[(req_obj.request_id, req_obj.index_in_shm_mem) for req_obj in req_objs]}" ) - req_status = ReqStatus(group_request_id, multimodal_params, req_objs, start_time) + req_status = ReqStatus( + group_request_id, + multimodal_params, + req_objs, + start_time, + decode_admission_lease=decode_admission_lease, + ) self.req_id_to_out_inf[group_request_id] = req_status + decode_admission_lease = None + unmanaged_req_indexes.clear() + unmanaged_req_objs.clear() # RL:请求已登记到 req_id_to_out_inf 并即将转发下游,从 admission gate # 注销,避免 pause 统计里仍把它算作“等待准入”的 pending 请求。 if self.rl_controller is not None: @@ -517,6 +553,10 @@ async def generate( # 已经放入到 req_id_to_out_inf 中的请求对象,由统一的回收循环 # 进行回收。 if group_request_id not in self.req_id_to_out_inf: + for req_obj in reversed(unmanaged_req_objs): + await self.shm_req_manager.async_put_back_req_obj(req_obj) + for req_index in reversed(unmanaged_req_indexes): + await self.shm_req_manager.async_release_req_index(req_index) await self._release_multimodal_resources(multimodal_params) await self.abort(group_request_id) raise e @@ -527,8 +567,25 @@ async def generate( await self.rl_controller.unregister_generation_admission(group_request_id) if running_request_registered: await self._unregister_running_request() + if decode_admission_lease is not None: + decode_admission_lease.release() return + async def _alloc_req_indexes(self, req_num: int, expect_available: bool) -> List[int]: + """Wait for one atomic batch unless Decode admission already reserved it.""" + indexes = await self.shm_req_manager.async_alloc_req_indexes(req_num) + if indexes is not None: + return indexes + if expect_available: + raise RuntimeError("Decode admission reserved slots that are unavailable in shm_req") + + sleep_time = 0.1 + while indexes is None: + await asyncio.sleep(sleep_time) + sleep_time = min(1, sleep_time * 1.1) + indexes = await self.shm_req_manager.async_alloc_req_indexes(req_num) + return indexes + def _count_multimodal_tokens(self, multimodal_params: MultimodalParams) -> Tuple[int, int]: image_tokens = 0 audio_tokens = 0 @@ -900,6 +957,7 @@ async def recycle_resource_loop(self): await self.shm_req_manager.async_put_back_req_obj(req) await self.shm_req_manager.async_release_req_index(req.index_in_shm_mem) await self._release_multimodal_resources(req_status.group_req_objs.multimodal_params) + req_status.release_decode_admission() # 先保留这个关键得日志,用于方便定位重构中的问题。 if time.time() - pre_time_mark > 120: @@ -925,7 +983,10 @@ async def handle_loop(self): if self.is_multinode_tp_slave: asyncio.create_task(self.loop_for_request()) - if self.pd_mode.is_P_or_D(): + # Multi-node TP rank 0 is the single HTTP/PD ingress and forwards each + # admitted request to its slave ranks. Slaves must not register as + # independent P/D nodes or they could bypass rank 0 admission. + if self.pd_mode.is_P_or_D() and not self.is_multinode_tp_slave: from lightllm.server.httpserver.pd_loop import pd_handle_loop asyncio.create_task(pd_handle_loop(self)) @@ -1032,7 +1093,14 @@ async def _unregister_running_request(self): class ReqStatus: - def __init__(self, group_request_id, multimodal_params, req_objs: List[Req], start_time) -> None: + def __init__( + self, + group_request_id, + multimodal_params, + req_objs: List[Req], + start_time, + decode_admission_lease: Optional[DecodeAdmissionLease] = None, + ) -> None: self.lock = asyncio.Lock() self.event = asyncio.Event() self.group_req_objs = GroupReqObjs( @@ -1042,6 +1110,13 @@ def __init__(self, group_request_id, multimodal_params, req_objs: List[Req], sta time_mark=start_time, ) self.out_token_info_list = [] + self.decode_admission_lease = decode_admission_lease + + def release_decode_admission(self) -> None: + if self.decode_admission_lease is None: + return + self.decode_admission_lease.release() + self.decode_admission_lease = None def can_release(self): for req in self.group_req_objs.shm_req_objs: diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index 95a280bdf4..bc21856c7a 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -9,66 +9,30 @@ 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, - PD_MASTER_CAPACITY_EPOCH_KEY, - PD_MASTER_CAPACITY_SHARE_KEY, -) +from lightllm.server.pd_io_struct import PD_DECODE_ADMISSION_CAPABILITY_KEY, 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.server.httpserver.manager import HttpServerManager +from lightllm.server.httpserver.decode_admission import DecodeAdmissionLeaseHandle 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.error_utils import PDPrefillNodeStopGenToken, ServerBusyError from lightllm.utils.shm_port_args import get_shm_port_args logger = init_logger(__name__) -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 - 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 _build_pd_registration_info(manager: HttpServerManager, pd_master_obj: PD_Master_Obj) -> dict: +def _build_pd_registration_info(manager: HttpServerManager) -> 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 + if manager.pd_mode.is_D(): + args_dict[PD_DECODE_ADMISSION_CAPABILITY_KEY] = manager.decode_admission_controller is not None return { "node_id": manager.args.pd_node_id, "client_ip_port": f"{manager.host_ip}:{get_shm_port_args().port}", @@ -85,6 +49,52 @@ async def timer_log(manager: HttpServerManager): return +async def _reserve_decode_slots( + manager: HttpServerManager, + reservation_id: int, + request_ids: tuple[int, ...], + reserved_lease_handles: Dict[int, DecodeAdmissionLeaseHandle], + websocket: ClientConnection, +) -> None: + installed_handles: Dict[int, DecodeAdmissionLeaseHandle] = {} + completed = False + lease = None + try: + controller = manager.decode_admission_controller + if controller is None: + raise ServerBusyError("Decode node admission is unavailable") + + lease = await controller.acquire(len(request_ids)) + manager.metric_client.histogram_observe("lightllm_request_queue_duration_bucket", lease.waited_seconds) + child_leases = lease.split([1] * len(request_ids)) + lease = None + installed_handles = { + request_id: DecodeAdmissionLeaseHandle(child_lease) + for request_id, child_lease in zip(request_ids, child_leases) + } + reserved_lease_handles.update(installed_handles) + await websocket.send(pickle.dumps((ObjType.PD_DECODE_SLOTS_RESERVED, reservation_id))) + completed = True + except asyncio.CancelledError: + raise + except ServerBusyError as error: + logger.warning(f"Decode reservation {reservation_id} rejected: {error.message}") + await websocket.send(pickle.dumps((ObjType.PD_UPLOAD_SERVER_BUSY, reservation_id, error.message))) + except BaseException as error: + logger.exception(f"Decode reservation {reservation_id} failed: {str(error)}") + await websocket.send( + pickle.dumps((ObjType.PD_UPLOAD_SERVER_BUSY, reservation_id, f"{type(error).__name__}: {str(error)}")) + ) + finally: + if not completed: + if lease is not None: + lease.release() + for request_id, handle in installed_handles.items(): + if reserved_lease_handles.get(request_id) is handle: + reserved_lease_handles.pop(request_id, None) + handle.release() + + async def pd_handle_loop(manager: HttpServerManager): if manager.args.host in ["127.0.0.1", "localhost"]: logger.error("pd mode must specify host ip, not use 127.0.0.1 or localhost") @@ -107,7 +117,6 @@ 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: - _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() @@ -137,6 +146,9 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O forwarding_tokens_task = None heartbeat_task = None generation_tasks: Dict[int, asyncio.Task] = {} + reservation_tasks: Dict[int, asyncio.Task] = {} + request_id_to_reservation_task: Dict[int, asyncio.Task] = {} + reserved_lease_handles: Dict[int, DecodeAdmissionLeaseHandle] = {} try: uri = f"ws://{pd_master_obj.host_ip_port}/pd_register" async with websockets.connect( @@ -146,32 +158,78 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O # 下方应用层心跳已负责存活检测,禁用协议层 keepalive,避免繁忙连接被误断。 ping_interval=None, ) as websocket: - sock = websocket.transport.get_extra_info("socket") sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) # 发送注册信息 - regist_json = _build_pd_registration_info(manager, pd_master_obj) + regist_json = _build_pd_registration_info(manager) 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, pd_master_obj.node_id) - ) - heartbeat_task = asyncio.create_task( - _send_heartbeat_to_pd_master(manager, websocket, pd_master_obj.node_id) - ) + 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)) group_req_id_to_event: Dict[int, asyncio.Event] = weakref.WeakValueDictionary() # 接收 pd master 发来的请求,并推理后,将生成的token转发回pd master。 while True: recv_bytes = await websocket.recv() obj = pickle.loads(recv_bytes) - if obj[0] == ObjType.REQ: + if obj[0] == ObjType.PD_RESERVE_DECODE_SLOTS: + _, reservation_id, request_ids = obj + request_ids = tuple(request_ids) + if ( + not request_ids + or reservation_id in reservation_tasks + or len(set(request_ids)) != len(request_ids) + or any( + request_id in request_id_to_reservation_task + or request_id in reserved_lease_handles + or request_id in generation_tasks + for request_id in request_ids + ) + ): + await websocket.send( + pickle.dumps( + ( + ObjType.PD_UPLOAD_SERVER_BUSY, + reservation_id, + "invalid or duplicate Decode reservation", + ) + ) + ) + continue + + reservation_task = asyncio.create_task( + _reserve_decode_slots( + manager, reservation_id, request_ids, reserved_lease_handles, websocket + ) + ) + reservation_tasks[reservation_id] = reservation_task + for request_id in request_ids: + request_id_to_reservation_task[request_id] = reservation_task + + def remove_reservation_task( + task: asyncio.Task, + current_reservation_id: int = reservation_id, + current_request_ids: tuple[int, ...] = request_ids, + ): + if reservation_tasks.get(current_reservation_id) is task: + reservation_tasks.pop(current_reservation_id, None) + for request_id in current_request_ids: + if request_id_to_reservation_task.get(request_id) is task: + request_id_to_reservation_task.pop(request_id, None) + if not task.cancelled(): + error = task.exception() + if error is not None: + logger.error(f"Decode reservation task failed: {str(error)}") + + reservation_task.add_done_callback(remove_reservation_task) + elif obj[0] == ObjType.REQ: prompt, sampling_params, multimodal_params = obj[1] group_req_id = sampling_params.group_request_id + decode_admission_lease_handle = reserved_lease_handles.pop(group_req_id, None) pd_event = asyncio.Event() group_req_id_to_event[group_req_id] = pd_event generation_task = asyncio.create_task( @@ -183,6 +241,7 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O forwarding_queue=forwarding_queue, pd_upload_websocket=websocket, pd_event=pd_event, + decode_admission_lease_handle=decode_admission_lease_handle, ) ) generation_tasks[group_req_id] = generation_task @@ -195,6 +254,12 @@ def remove_generation_task(task: asyncio.Task, request_id: int = group_req_id): elif obj[0] == ObjType.ABORT: group_req_id = obj[1] logger.warning(f"recv cmd aborted req id {group_req_id}") + reservation_task = request_id_to_reservation_task.get(group_req_id) + if reservation_task is not None and not reservation_task.done(): + reservation_task.cancel() + reserved_handle = reserved_lease_handles.pop(group_req_id, None) + if reserved_handle is not None: + reserved_handle.release() generation_task = generation_tasks.get(group_req_id) if generation_task is not None and not generation_task.done(): generation_task.cancel() @@ -230,10 +295,14 @@ async def delayed_abort_task(group_req_id, retry_count): finally: child_tasks = [task for task in (forwarding_tokens_task, heartbeat_task) if task is not None] child_tasks.extend(generation_tasks.values()) + child_tasks.extend(reservation_tasks.values()) for task in child_tasks: task.cancel() if child_tasks: await asyncio.gather(*child_tasks, return_exceptions=True) + for handle in reserved_lease_handles.values(): + handle.release() + reserved_lease_handles.clear() await asyncio.sleep(10) await forwarding_queue.get_all_data() @@ -282,6 +351,7 @@ async def _pd_process_generate( forwarding_queue: AsyncQueue, pd_upload_websocket: ClientConnection, pd_event: asyncio.Event, + decode_admission_lease_handle: Optional[DecodeAdmissionLeaseHandle] = None, ): try: async for sub_req_id, request_output, metadata, finish_status in manager.generate( @@ -291,11 +361,19 @@ async def _pd_process_generate( request=None, pd_upload_websocket=pd_upload_websocket, pd_event=pd_event, + decode_admission_lease_handle=decode_admission_lease_handle, ): metadata["node_mode"] = manager.args.run_mode await forwarding_queue.put((sub_req_id, request_output, metadata, finish_status)) except PDPrefillNodeStopGenToken as e: logger.info(f"pd prefill node stop gen token for group_request_id {e.group_request_id}") + except ServerBusyError as e: + group_request_id = sampling_params.group_request_id + logger.warning(f"pd node rejected request {group_request_id}: {e.message}") + try: + await pd_upload_websocket.send(pickle.dumps((ObjType.PD_UPLOAD_SERVER_BUSY, group_request_id, e.message))) + except Exception: + logger.exception(f"report pd node request rejection failed, group_request_id: {group_request_id}") except asyncio.CancelledError: # PD master 主动 abort 或连接断开清理任务时会走取消路径,不需要反向重复上报。 pass @@ -310,44 +388,32 @@ async def _pd_process_generate( ) except Exception: logger.exception(f"report pd node generate error failed, group_request_id: {group_request_id}") + finally: + if decode_admission_lease_handle is not None: + decode_admission_lease_handle.release() # 转发token的task -async def _up_tokens_to_pd_master( - forwarding_queue: AsyncQueue, - websocket: ClientConnection, - pd_master_node_id: int, -): +async def _up_tokens_to_pd_master(forwarding_queue: AsyncQueue, websocket: ClientConnection): """批量向 PD Master 转发生成结果和最新负载。""" while True: handle_list = await forwarding_queue.wait_to_get_all_data() if handle_list: - load_info: dict = _get_load_info(pd_master_node_id) + load_info: dict = _get_load_info() await websocket.send(pickle.dumps((ObjType.TOKEN_PACKS, handle_list, load_info))) -async def _send_heartbeat_to_pd_master( - manager: HttpServerManager, - websocket: ClientConnection, - pd_master_node_id: int, -): - """定时或在成员变化时向 PD Master 上报心跳。""" +async def _send_heartbeat_to_pd_master(websocket: ClientConnection): heartbeat_interval_seconds = 15 - membership_changed = manager.pd_master_membership_changed while True: - 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 + await websocket.send(pickle.dumps((ObjType.HEARTBEAT,))) + await asyncio.sleep(heartbeat_interval_seconds) # 获取节点负载信息 -def _get_load_info(pd_master_node_id: int) -> dict: - """汇总当前 Master 对应的容量和节点负载。""" +def _get_load_info() -> dict: + """汇总当前节点负载。""" from lightllm.server.api_http import g_objs @@ -361,11 +427,8 @@ 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) - 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), } return load_info diff --git a/lightllm/server/httpserver_for_pd_master/admission.py b/lightllm/server/httpserver_for_pd_master/admission.py deleted file mode 100644 index bebd026869..0000000000 --- a/lightllm/server/httpserver_for_pd_master/admission.py +++ /dev/null @@ -1,716 +0,0 @@ -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 - - def __post_init__(self) -> None: - """校验请求需要原子获取的 Decode 槽位数。""" - if self.decode_slots < 1: - raise ValueError("decode_slots must be positive") - - -class SessionTracker: - """只把服务端已经成功观察过的 Session 视为连续会话。""" - - def __init__( - self, - ttl_seconds: float, - 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() - 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: - """记录 Session 最近一次成功返回结果的时间。""" - 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: - request: AdmissionRequest - enqueue_time: float - deadline: float - deadline_changed: asyncio.Event - future: asyncio.Future - - -class AdmissionLease: - """一次已经获得的 Decode 容量租约。""" - - def __init__( - self, - controller: "PDAdmissionController", - request: AdmissionRequest, - waited_seconds: float, - ) -> None: - """保存本次租约占用的 Decode 槽位。""" - 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], - 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._clock = clock - self._state_change_callback = state_change_callback - self._active_slots = 0 - 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._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 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())) - - 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, 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( - request=request, - enqueue_time=enqueue_time, - deadline=enqueue_time + self.policy.max_wait_seconds(request.priority), - deadline_changed=asyncio.Event(), - future=loop.create_future(), - ) - - if not self._make_queue_room(waiter): - raise ServerBusyError("PD master admission queue is full") - - self._enqueue(waiter) - self._drain() - - try: - 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: - 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: - """Decode 容量变化后重置临时公平状态并重新驱动队列。""" - self._clear_backfill_state() - self._reset_deficits() - 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 - - 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: - 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) - if not self._queues[old_priority]: - self._deficits[old_priority] = 0 - self._drain() - - 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 - 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 self._active_slots < 0: - raise RuntimeError("PD admission active 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: - """把等待项加入优先级队列和 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: - 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) - 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]: - """取消等待项,或取回已并发发放的租约用于释放。""" - 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 - - 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 串行和 FIFO 约束。""" - 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, - excluded: Optional[_WaitingRequest] = None, - available_slots: Optional[int] = None, - ) -> Optional[_WaitingRequest]: - """按类内 FIFO 返回第一个满足 Session 和可选槽位约束的等待项。""" - for waiter in self._queues[priority]: - if waiter is excluded: - continue - if not self._session_is_grantable(waiter): - continue - if available_slots is not None and waiter.request.decode_slots > available_slots: - continue - 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 - - 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 - - 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 - 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 - - 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 - - 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 - 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 - - 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 - self._grant_waiter(waiter) - 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 fba20821d4..ae2f18e93a 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -4,7 +4,6 @@ import uvloop import time import datetime -import math import ujson as json import pickle import httpx @@ -14,12 +13,11 @@ from typing import Union, List, Tuple, Dict, Optional from lightllm.server.core.objs import FinishStatus from ..pd_io_struct import ( + PD_DECODE_ADMISSION_CAPABILITY_KEY, 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 @@ -34,17 +32,64 @@ 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, - PDAdmissionController, - SessionTracker, -) logger = init_logger(__name__) +class _DecodeReservationStatus: + def __init__(self) -> None: + self.event = asyncio.Event() + self.error_info: Optional[str] = None + + def raise_if_error(self) -> None: + if self.error_info is not None: + raise ServerBusyError(self.error_info) + + +class _DecodeAdmissionGroup: + """Wait for every choice's Prefill before reserving Decode slots once.""" + + def __init__( + self, + manager: "HttpServerManagerForPDMaster", + d_node: PD_Client_Obj, + reservation_id: int, + expected_count: int, + request: Request, + ) -> None: + self.manager = manager + self.d_node = d_node + self.reservation_id = reservation_id + self.expected_count = expected_count + self.request = request + self._lock = asyncio.Lock() + self._ready = asyncio.Event() + self._request_ids: List[int] = [] + self._error: Optional[BaseException] = None + + async def wait(self, group_request_id: int) -> None: + is_last = False + async with self._lock: + if group_request_id in self._request_ids: + raise RuntimeError(f"duplicate Decode reservation request id {group_request_id}") + self._request_ids.append(group_request_id) + is_last = len(self._request_ids) == self.expected_count + + if is_last: + try: + await self.manager.reserve_decode_slots( + self.d_node, self.reservation_id, tuple(self._request_ids), self.request + ) + except BaseException as error: + self._error = error + finally: + self._ready.set() + + await self._ready.wait() + if self._error is not None: + raise self._error + + class HttpServerManagerForPDMaster: def __init__( self, @@ -58,21 +103,8 @@ 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._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, - policy=self.admission_policy, - state_change_callback=self._record_admission_state, - ) - self.req_id_to_out_inf: Dict[int, ReqStatus] = {} + self.decode_reservation_statuses: Dict[int, _DecodeReservationStatus] = {} self.infos_queues = None # 这个需要延迟初始化,否则使用的loop不对 self.health_timeout = int(os.getenv("HEALTH_TIMEOUT", "200")) self.latest_success_infer_time = time.time() @@ -104,12 +136,10 @@ 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): @@ -122,44 +152,6 @@ async def update_req_status(self, upkv_status: PDUpKVStatus): pass return - def update_node_load_info(self, load_info: Optional[dict]) -> None: - """更新节点遥测;仅 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: - """合并同一事件循环周期内的状态变化,避免重复发送 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 prompt_ids = self.tokenizer.encode(prompt, None, **kwargs) @@ -198,69 +190,16 @@ async def generate( multimodal_params: MultimodalParams, request: Request, ): - admission_lease = None - running_request_registered = False - session_key = self._get_session_key(request) + 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 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: - if running_request_registered: - self.running_request_count -= 1 - if admission_lease is not None: - 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() - return session_key or None - - def _build_admission_request( - self, - prompt: Union[str, List[int]], - 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 - 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)) - return AdmissionRequest( - session_key=session_key, - priority=priority, - decode_slots=decode_slots, - ) + self.running_request_count -= 1 async def _generate( self, @@ -285,14 +224,33 @@ async def _generate( origin_sampling_params = SamplingParams.from_buffer_copy(sampling_params) origin_group_request_id = self.id_gen.generate_id() - # Record one user request even when it is expanded into multiple independent - # n=1 requests below. The externally visible ids remain in the same request - # group so OpenAI streaming can derive choice_index from the sub request id. + # Record one user request even when it is expanded into independent n=1 + # requests. The Decode reservation coordinator still accounts for n as + # one atomic unit after all Prefill stages are ready. await self._log_req_header(request, origin_group_request_id) self.metric_client.counter_inc("lightllm_request_count") self.metric_client.histogram_observe("lightllm_request_max_new_tokens", origin_sampling_params.max_new_tokens) - choice_count = origin_sampling_params.n + p_node, d_node = await self.select_p_d_node(prompt, origin_sampling_params, multimodal_params) + if not p_node or not d_node: + logger.error(f"{origin_group_request_id}: No p_node or d_node found") + raise Exception(f"{origin_group_request_id}: No p_node or d_node found") + + choice_count = max(1, int(origin_sampling_params.n or 1)) + decode_reservation = None + args = getattr(self, "args", None) + decode_admission_disabled = args is not None and ( + args.disable_pd_node_decode_admission or args.disable_pd_master_decode_capacity_limit + ) + if choice_count > 1 and not decode_admission_disabled: + decode_reservation = _DecodeAdmissionGroup( + manager=self, + d_node=d_node, + reservation_id=origin_group_request_id, + expected_count=choice_count, + request=request, + ) + generators = [] for choice_index in range(choice_count): choice_sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params) @@ -306,6 +264,9 @@ async def _generate( request, start_time, origin_group_request_id + choice_index, + p_node=p_node, + d_node=d_node, + decode_reservation=decode_reservation, ) ) @@ -322,25 +283,27 @@ async def _generate_one( request: Request, start_time: float, origin_request_id: int, + p_node: PD_Client_Obj, + d_node: PD_Client_Obj, + decode_reservation: Optional["_DecodeAdmissionGroup"] = None, ): # 先将请求根据max_new_tokens 参数进行分块操作,主要是 pd 分离场景中, # 只能使用保守调度,但是如果用户都设置一个很大的 max_new_tokens 值,会 # 导致极大显存预留,照成系统的吞吐能力下降,所以我们将请求分割成几段进行 # 推理,只要保证分块合理,实际分段推理是极少发生的情况,系统吞吐就不会受 # 到影响。 - max_new_tokens_list = self._split_max_new_tokens(max_new_tokens=origin_sampling_params.max_new_tokens) + if decode_reservation is None: + max_new_tokens_list = self._split_max_new_tokens(max_new_tokens=origin_sampling_params.max_new_tokens) + else: + # Choices have independent continuation histories, so the existing + # segmented continuation protocol cannot safely splice an n-way + # request. Each choice remains one n=1 request. + max_new_tokens_list = [origin_sampling_params.max_new_tokens] block_group_request_id = origin_request_id - p_node = None - d_node = None pending_prefill_load_chars = None try: - p_node, d_node = await self.select_p_d_node(prompt, origin_sampling_params, multimodal_params) - if not p_node or not d_node: - logger.error(f"{origin_request_id}: No p_node or d_node found") - raise Exception(f"{origin_request_id}: No p_node or d_node found") - history_gen_token_strs = [] origin_prompt_cache_len = None @@ -365,11 +328,11 @@ async def _generate_one( sampling_params, multimodal_params, request, + decode_reservation=decode_reservation, ) is_last_block = iter_index == len(max_new_tokens_list) - 1 prompt_tokens = sys.maxsize # 因为分段的原因 async for sub_req_id, request_output, metadata, finish_status in results_generator: - # pd 分离模式下,返回的 metadata 可能序号信息可能存在不准确性。 assert sub_req_id == block_group_request_id if finish_status.is_finished_length() and not is_last_block: finish_status = FinishStatus() # 转换为NoFinished @@ -476,6 +439,24 @@ async def raise_if_disconnected() -> None: except asyncio.TimeoutError: continue + async def reserve_decode_slots( + self, d_node: PD_Client_Obj, reservation_id: int, request_ids: Tuple[int, ...], request: Request + ) -> None: + status = _DecodeReservationStatus() + self.decode_reservation_statuses[reservation_id] = status + try: + await d_node.websocket.send_bytes( + pickle.dumps((ObjType.PD_RESERVE_DECODE_SLOTS, reservation_id, request_ids)) + ) + timeout = float(d_node.start_args.get("pd_node_decode_admission_timeout", 5)) + 5 + await self._wait_for_event_or_disconnect( + status.event, request, timeout=timeout, group_request_id=reservation_id, stage="decode admission" + ) + status.raise_if_error() + finally: + if self.decode_reservation_statuses.get(reservation_id) is status: + self.decode_reservation_statuses.pop(reservation_id, None) + async def _log_req_header(self, request: Request, group_request_id: int): x_request_id = request.headers.get("X-Request-Id", "") x_session_id = request.headers.get("X-Session-Id", "") @@ -495,6 +476,7 @@ async def fetch_pd_stream( sampling_params: SamplingParams, multimodal_params: MultimodalParams, request: Request, + decode_reservation: Optional["_DecodeAdmissionGroup"] = None, ): group_request_id = sampling_params.group_request_id sampling_params.pd_master_node_id.initialize(self.args.pd_node_id) @@ -525,6 +507,9 @@ async def fetch_pd_stream( prompt_ids = prefill_prompt_ids_event.prompt_ids logger.info(f"group_request_id: {group_request_id} get prefill prompt ids len {len(prompt_ids)}") + if decode_reservation is not None: + await decode_reservation.wait(group_request_id) + sampling_params.max_new_tokens = old_max_new_tokens await d_node.websocket.send_bytes( pickle.dumps((ObjType.REQ, (prompt_ids, sampling_params, MultimodalParams()))) @@ -632,6 +617,7 @@ async def _wait_to_token_package( sampling_params: SamplingParams, multimodal_params: MultimodalParams, request: Request, + decode_reservation: Optional["_DecodeAdmissionGroup"] = None, ): if sampling_params.disable_prompt_cache: assert False, "pd mode dont support set disable_prompt_cache to True" @@ -646,7 +632,13 @@ async def _wait_to_token_package( sub_req_id_to_mtp_verify_step_num: Dict[int, int] = {} async for sub_req_id, out_str, metadata, finish_status in self.fetch_pd_stream( - p_node, d_node, prompt, sampling_params, multimodal_params, request + p_node, + d_node, + prompt, + sampling_params, + multimodal_params, + request, + decode_reservation=decode_reservation, ): if await request.is_disconnected(): raise ClientDisconnected( @@ -759,7 +751,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.update_node_load_info(node_load_info) + self.pd_manager.update_node_load_info(node_load_info) for sub_req_id, text, metadata, finish_status in token_list: finish_status: FinishStatus = finish_status @@ -794,6 +786,28 @@ async def handle_loop(self): ) else: await req_status.set_error(error_info) + elif obj[0] == ObjType.PD_UPLOAD_SERVER_BUSY: + _, group_req_id, error_info = obj + logger.warning( + f"received PD node server busy, group_req_id: {group_req_id}, reason: {error_info}" + ) + reservation_status = getattr(self, "decode_reservation_statuses", {}).get(group_req_id) + if reservation_status is not None: + reservation_status.error_info = error_info + reservation_status.event.set() + continue + req_status = self.req_id_to_out_inf.get(group_req_id) + if req_status is None: + logger.error(f"PD_UPLOAD_SERVER_BUSY fail find req status for group_req_id: {group_req_id}") + else: + await req_status.set_error(error_info, is_server_busy=True) + elif obj[0] == ObjType.PD_DECODE_SLOTS_RESERVED: + _, reservation_id = obj + reservation_status = getattr(self, "decode_reservation_statuses", {}).get(reservation_id) + if reservation_status is None: + logger.warning(f"received stale Decode reservation ack: {reservation_id}") + else: + reservation_status.event.set() else: logger.error(f"recevie error obj {obj}") except BaseException as e: @@ -820,6 +834,7 @@ def __init__(self, req_id, p_node, d_node) -> None: self.p_node: PD_Client_Obj = p_node self.d_node: PD_Client_Obj = d_node self.error_info: Optional[str] = None + self.is_server_busy = False async def wait_to_ready(self): try: @@ -827,9 +842,10 @@ async def wait_to_ready(self): except asyncio.TimeoutError: pass - async def set_error(self, error_info: str): + async def set_error(self, error_info: str, is_server_busy: bool = False): async with self.lock: self.error_info = error_info + self.is_server_busy = is_server_busy # 请求可能正在等待 Prefill prompt ids、Decode KV 资源或输出 token, # 设置全部事件,让请求自己的执行循环立即醒来并抛出异常。 self.event.set() @@ -838,6 +854,8 @@ async def set_error(self, error_info: str): def raise_if_error(self): if self.error_info is not None: + if self.is_server_busy: + raise ServerBusyError(self.error_info) logger.error( f"group_request_id: {self.req_id} detected PD node generate error, " f"raise exception to end the request flow early: {self.error_info}" @@ -879,13 +897,6 @@ def __init__(self, args: StartArgs): self.selector = create_selector(args.select_p_d_node_strategy, self) 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 - ) - def is_pd_nodes_ready(self): prefill_node_count = len(self.prefill_nodes) decode_node_count = len(self.decode_nodes) @@ -937,24 +948,24 @@ async def check_pd_nodes_health(self): return True def register_pd(self, pd_info_json, websocket): - # 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. + # Ignore the short-lived top-level capacity fields during a rolling + # upgrade from the former Master-side admission implementation. pd_info = dict(pd_info_json) + pd_info.pop("capacity_share", None) + pd_info.pop("capacity_epoch", None) 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, + admission_disabled = ( + self.args.disable_pd_node_decode_admission or self.args.disable_pd_master_decode_capacity_limit ) + if ( + pd_info.get("mode") == "decode" + and not admission_disabled + and start_args.get(PD_DECODE_ADMISSION_CAPABILITY_KEY) is not True + ): + raise ValueError( + "Decode node does not provide authoritative admission; upgrade Decode nodes before PD Masters" + ) + pd_client = PD_Client_Obj(**pd_info) 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( @@ -994,25 +1005,22 @@ def register_pd(self, pd_info_json, websocket): return def remove_pd(self, 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") + pd_info = dict(pd_info_json) + pd_info.pop("capacity_share", None) + pd_info.pop("capacity_epoch", None) + pd_client = PD_Client_Obj(**pd_info) - 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.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.selector.update_nodes(self.prefill_nodes, self.decode_nodes) - logger.info(f"mode: {mode} url: {client_ip_port} removed") + logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} removed") return - def update_node_load_info(self, load_info: Optional[dict]) -> bool: - """更新节点负载信息,并返回有效 Decode 容量份额是否发生变化。 - - capacity_epoch 只用于拒绝旧上报;仅 epoch 变化不会重新驱动 admission。 - + def update_node_load_info(self, load_info: Optional[dict]): + """更新节点负载信息 load_info: 节点负载信息字典,内容格式如下,可以为 None { "total_token_usage_rate": xxxx, @@ -1021,28 +1029,14 @@ def update_node_load_info(self, load_info: Optional[dict]) -> bool: """ try: if load_info is None: - return False + 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) - if pd_client is None: - return False - pd_client.run_status.total_token_usage_rate = load_info["total_token_usage_rate"] - - 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"] - ) - 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: + pd_client.run_status.total_token_usage_rate = total_token_usage_rate + except BaseException as e: logger.warning(f"udpate node load info failed, load_info: {load_info} error: {str(e)}") - return False + return def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams 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 6bb9bdf6ef..adf911bc21 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,14 +19,13 @@ 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 PromptCacheMatchResult, PromptCacheTree +from .prompt_cache_tree import PromptCacheTree logger = init_logger(__name__) @@ -57,18 +56,6 @@ 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 命中率动态调整负载均衡阈值。""" @@ -169,36 +156,9 @@ 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: - """估算可复用缓存的命中率,并在当前异步请求上下文中保留匹配结果。""" - 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() - # 准入和选点共享同一个提示词对象。通过对象身份判断可以避免比较可能很长的字符串, - # 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 - _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._match_prompt_cache(request_text) + result = self.prompt_cache_tree.prefix_match(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/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py b/lightllm/server/httpserver_for_pd_master/pd_selector/pd_selector.py index eb0f9b5704..5474806b7b 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, Optional +from typing import Union, List, Tuple, Dict 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,10 +30,6 @@ 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]: - """当选择器无法估算可复用的提示词缓存时返回 None。""" - return None - class RandomSelector(PDSelector): """随机选择器""" @@ -104,9 +100,3 @@ 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]: - """返回 cache-aware 策略对当前提示词的命中估算。""" - if not isinstance(prompt, str): - return 0.0 - return self.policy.estimate_cache_hit_rate(self.prefill_nodes, prompt) diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index 5814b11f61..0d42462c3f 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -32,9 +32,6 @@ "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_decode_capacity": "Decode slot capacity currently assigned to the PD master", } @@ -114,9 +111,6 @@ 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_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 f60a470c6e..30ba92ca27 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -10,12 +10,7 @@ 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" +PD_DECODE_ADMISSION_CAPABILITY_KEY = "__pd_decode_admission_v1" # 节点的行为 @@ -49,6 +44,9 @@ class ObjType(enum.Enum): PD_REQ_DECODE_NODE_INFO = 5 # pd master 节点下发给 prefill 节点的请求对应的 decode 节点信息。 HEARTBEAT = 6 # P/D 节点向 pd master 上报的心跳。 PD_UPLOAD_GENERATE_ERROR = 7 # P/D 节点向 pd master 上报本地请求生成异常。 + PD_UPLOAD_SERVER_BUSY = 8 # P/D 节点向 pd master 上报本地准入拒绝。 + PD_RESERVE_DECODE_SLOTS = 9 # PD master 在 Decode 节点原子预留一组请求槽。 + PD_DECODE_SLOTS_RESERVED = 10 # Decode 节点确认一组请求槽已预留。 @dataclass @@ -64,9 +62,6 @@ 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 的请求数。 @@ -77,10 +72,6 @@ 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): @@ -106,7 +97,6 @@ class PDUpKVStatus: pd_kv_trans_params: bytes # pd kv 传输建立连接所使用的元数据对象 def __post_init__(self): - if not isinstance(self.group_request_id, int): error_info = "group_request_id only can be int" logger.error(error_info) diff --git a/test/test_pd_selector/test_pd_master_multi_choice.py b/test/test_pd_selector/test_pd_master_multi_choice.py index 7d0b32ded9..12b52e787d 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -1,6 +1,8 @@ from __future__ import annotations import asyncio +import pickle +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, call, patch import pytest @@ -8,20 +10,29 @@ from lightllm.server.core.objs import FinishStatus, SamplingParams from lightllm.server.httpserver.manager import HttpServerManager from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster +from lightllm.server.multimodal_params import MultimodalParams +from lightllm.server.pd_io_struct import ObjType def _manager() -> HttpServerManagerForPDMaster: manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.pd_manager = MagicMock() + manager.args = SimpleNamespace( + disable_pd_node_decode_admission=False, + disable_pd_master_decode_capacity_limit=False, + pd_node_id=7, + ) manager.id_gen = MagicMock() manager.id_gen.generate_id.return_value = 800 manager.metric_client = MagicMock() manager._log_req_header = AsyncMock() manager.tokens = MagicMock(return_value=2) + manager.req_id_to_out_inf = {} + manager.decode_reservation_statuses = {} return manager -def test_pd_master_expands_n_into_concurrent_single_choice_requests(): +def test_pd_master_reserves_n_then_sends_independent_choices_to_one_decode_node(): async def run(): manager = _manager() manager.id_gen.generate_id.side_effect = [800, 808, 816, 824] @@ -33,38 +44,34 @@ async def run(): multimodal_params = MagicMock() multimodal_params.verify_and_preload = AsyncMock() request = MagicMock() - - started = set() - all_started = asyncio.Event() - captured_params = [] p_node = MagicMock(dispatched_prompt_chars=0, dispatched_req_num=0) d_node = MagicMock() manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) - manager._split_max_new_tokens = MagicMock(return_value=[4]) + manager._split_max_new_tokens = MagicMock() manager.remove_req = AsyncMock() + manager.reserve_decode_slots = AsyncMock() + captured_params = [] + captured_nodes = [] async def wait_to_token_package( selected_p_node, selected_d_node, start_time, prompt, - choice_sampling_params, + node_sampling_params, child_multimodal_params, child_request, + decode_reservation=None, ): - captured_params.append(choice_sampling_params) - started.add(choice_sampling_params.group_request_id) - if len(started) == 3: - all_started.set() - await all_started.wait() - for token_index in range(2): - finish_status = FinishStatus() if token_index == 0 else FinishStatus(FinishStatus.FINISHED_STOP) - yield ( - choice_sampling_params.group_request_id, - f"internal-{choice_sampling_params.group_request_id}-{token_index}", - {"prompt_tokens": 2}, - finish_status, - ) + captured_params.append(node_sampling_params) + captured_nodes.append((selected_p_node, selected_d_node)) + await decode_reservation.wait(node_sampling_params.group_request_id) + yield ( + node_sampling_params.group_request_id, + f"choice-{node_sampling_params.group_request_id}", + {"prompt_tokens": 2}, + FinishStatus(FinishStatus.FINISHED_STOP), + ) manager._wait_to_token_package = wait_to_token_package @@ -73,21 +80,21 @@ async def wait_to_token_package( async for result in manager._generate("prompt", sampling_params, multimodal_params, request): results.append(result) - assert [result[0] for result in results].count(800) == 2 - assert [result[0] for result in results].count(801) == 2 - assert [result[0] for result in results].count(802) == 2 - assert {result[1] for result in results} == { - "internal-808-0", - "internal-808-1", - "internal-816-0", - "internal-816-1", - "internal-824-0", - "internal-824-1", - } - assert all(params.n == 1 and params.best_of == 1 for params in captured_params) + assert {result[0] for result in results} == {800, 801, 802} + assert {result[1] for result in results} == {"choice-808", "choice-816", "choice-824"} + assert len(captured_params) == 3 + assert all(params.n == 1 for params in captured_params) + assert all(params.best_of == 1 for params in captured_params) assert {params.group_request_id for params in captured_params} == {808, 816, 824} - assert sampling_params.n == 3 - assert sampling_params.best_of == 3 + assert captured_nodes == [(p_node, d_node)] * 3 + manager.select_p_d_node.assert_awaited_once() + manager._split_max_new_tokens.assert_not_called() + manager.reserve_decode_slots.assert_awaited_once() + reserve_args = manager.reserve_decode_slots.await_args.args + assert reserve_args[0] is d_node + assert reserve_args[1] == 800 + assert set(reserve_args[2]) == {808, 816, 824} + assert reserve_args[3] is request multimodal_params.verify_and_preload.assert_awaited_once_with(request) manager._log_req_header.assert_awaited_once_with(request, 800) assert manager.metric_client.counter_inc.call_args_list == [ @@ -100,6 +107,57 @@ async def wait_to_token_package( asyncio.run(asyncio.wait_for(run(), timeout=2)) +def test_pd_master_waits_for_prefill_and_atomic_reservation_before_decode_request(): + async def run(): + manager = _manager() + sampling_params = SamplingParams() + sampling_params.group_request_id = 808 + sampling_params.max_new_tokens = 4 + p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) + d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) + request = SimpleNamespace(is_disconnected=AsyncMock(return_value=False)) + reservation_entered = asyncio.Event() + reservation_release = asyncio.Event() + + async def wait_for_reservation(group_request_id): + assert group_request_id == 808 + reservation_entered.set() + await reservation_release.wait() + + decode_reservation = SimpleNamespace(wait=wait_for_reservation) + generator = manager.fetch_pd_stream( + p_node, + d_node, + "prompt", + sampling_params, + MultimodalParams(), + request, + decode_reservation=decode_reservation, + ) + next_result = asyncio.create_task(generator.__anext__()) + try: + while 808 not in manager.req_id_to_out_inf: + await asyncio.sleep(0) + req_status = manager.req_id_to_out_inf[808] + req_status.prefill_prompt_ids_event.prompt_ids = [1, 2, 3] + req_status.prefill_prompt_ids_event.set() + + await asyncio.wait_for(reservation_entered.wait(), timeout=1) + d_node.websocket.send_bytes.assert_not_awaited() + + reservation_release.set() + while d_node.websocket.send_bytes.await_count == 0: + await asyncio.sleep(0) + decode_message = pickle.loads(d_node.websocket.send_bytes.await_args.args[0]) + assert decode_message[0] == ObjType.REQ + assert list(decode_message[1][0]) == [1, 2, 3] + finally: + next_result.cancel() + await asyncio.gather(next_result, return_exceptions=True) + + asyncio.run(asyncio.wait_for(run(), timeout=2)) + + def test_pd_master_n_one_uses_the_same_choice_merge_path(): async def run(): manager = _manager() @@ -111,6 +169,9 @@ async def run(): multimodal_params = MagicMock() multimodal_params.verify_and_preload = AsyncMock() request = MagicMock() + p_node = MagicMock() + d_node = MagicMock() + manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) merge_choice_generators = manager._merge_choice_generators manager._merge_choice_generators = MagicMock(side_effect=merge_choice_generators) @@ -121,6 +182,7 @@ async def generate_one( child_request, start_time, origin_request_id, + **_kwargs, ): yield ( origin_request_id, @@ -217,7 +279,6 @@ async def run(): manager.abort = AsyncMock() p_node = MagicMock(dispatched_prompt_chars=0, dispatched_req_num=0) d_node = MagicMock() - manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) async def failing_wait_to_token_package(*_args, **_kwargs): raise RuntimeError("generation failed") @@ -234,6 +295,8 @@ async def failing_wait_to_token_package(*_args, **_kwargs): MagicMock(), 0, 800, + p_node, + d_node, ): pass @@ -257,13 +320,14 @@ async def run(): dispatched_req_num=other_request_count, ) d_node = MagicMock() - manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) dispatched_nodes = [] dispatched_prompts = [] dispatched_loads = [] dispatched_req_counts = [] - async def wait_to_token_package(selected_p_node, _d_node, _start_time, block_prompt, sampling_params, *_args): + async def wait_to_token_package( + selected_p_node, _d_node, _start_time, block_prompt, sampling_params, *_args, **_kwargs + ): dispatched_nodes.append(selected_p_node) dispatched_prompts.append(block_prompt) dispatched_loads.append(selected_p_node.dispatched_prompt_chars) @@ -285,10 +349,11 @@ async def wait_to_token_package(selected_p_node, _d_node, _start_time, block_pro MagicMock(), 0, 800, + p_node, + d_node, ): results.append(result) - manager.select_p_d_node.assert_awaited_once() assert dispatched_nodes == [p_node, p_node] assert dispatched_prompts == ["prompt", "promptx"] assert dispatched_loads == [other_request_load + len("prompt"), other_request_load + len("promptx")] @@ -314,7 +379,6 @@ async def run(): dispatched_req_num=other_request_count, ) d_node = MagicMock() - manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node)) async def wait_to_token_package(*_args, **_kwargs): yield 808, "first", {"prompt_tokens": 1}, FinishStatus() @@ -328,6 +392,8 @@ async def wait_to_token_package(*_args, **_kwargs): MagicMock(), 0, 800, + p_node, + d_node, ) assert (await generator.__anext__())[1] == "first" diff --git a/unit_tests/server/core/objs/test_shm_req_manager.py b/unit_tests/server/core/objs/test_shm_req_manager.py index 56871b5466..59fba2c05f 100644 --- a/unit_tests/server/core/objs/test_shm_req_manager.py +++ b/unit_tests/server/core/objs/test_shm_req_manager.py @@ -1,6 +1,7 @@ import os import pytest import time +import numpy as np from unittest.mock import MagicMock from easydict import EasyDict @@ -81,6 +82,25 @@ def test_put_back_req_obj(shm_req_manager): shm_req_manager.release_req_index(index) +def test_alloc_req_indexes_allocates_all_slots_together(shm_req_manager): + indexes = shm_req_manager.alloc_req_indexes(3) + + assert indexes is not None + assert len(indexes) == 3 + assert all(shm_req_manager.alloc_state_shm.arr[index] == 1 for index in indexes) + + for index in indexes: + shm_req_manager.release_req_index(index) + + +def test_alloc_req_indexes_rolls_back_when_full_batch_is_unavailable(shm_req_manager): + before = shm_req_manager.alloc_state_shm.arr.copy() + free_slots = int(np.sum(before == 0)) + + assert shm_req_manager.alloc_req_indexes(free_slots + 1) is None + np.testing.assert_array_equal(shm_req_manager.alloc_state_shm.arr, before) + + def test_alloc_req_index_no_available(shm_req_manager): for _ in range(shm_req_manager.max_req_num): shm_req_manager.alloc_req_index() diff --git a/unit_tests/server/httpserver/test_decode_admission.py b/unit_tests/server/httpserver/test_decode_admission.py new file mode 100644 index 0000000000..6883e88192 --- /dev/null +++ b/unit_tests/server/httpserver/test_decode_admission.py @@ -0,0 +1,169 @@ +import asyncio + +import pytest + +from lightllm.server.httpserver.decode_admission import ( + DecodeAdmissionController, + DecodeAdmissionLeaseHandle, +) +from lightllm.utils.error_utils import ServerBusyError + + +def test_decode_admission_waits_and_grants_in_fifo_order(): + async def run(): + controller = DecodeAdmissionController(capacity=2, max_queued_slots=2, timeout_seconds=1) + active = await controller.acquire(2) + first_task = asyncio.create_task(controller.acquire(1)) + second_task = asyncio.create_task(controller.acquire(1)) + await asyncio.sleep(0) + + assert controller.active_slots == 2 + assert controller.queued_slots == 2 + assert not first_task.done() + assert not second_task.done() + + active.release() + first = await first_task + second = await second_task + assert controller.active_slots == 2 + + first.release() + second.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_decode_admission_keeps_strict_fifo_when_head_gang_does_not_fit(): + async def run(): + controller = DecodeAdmissionController(capacity=3, max_queued_slots=3, timeout_seconds=1) + first_active = await controller.acquire(1) + second_active = await controller.acquire(1) + third_active = await controller.acquire(1) + gang_task = asyncio.create_task(controller.acquire(2)) + await asyncio.sleep(0) + small_task = asyncio.create_task(controller.acquire(1)) + await asyncio.sleep(0) + + first_active.release() + await asyncio.sleep(0) + assert not gang_task.done() + assert not small_task.done() + assert controller.active_slots == 2 + + second_active.release() + gang = await gang_task + assert not small_task.done() + third_active.release() + small = await small_task + assert controller.active_slots == 3 + + gang.release() + small.release() + + asyncio.run(run()) + + +def test_decode_admission_cancellation_removes_waiter(): + async def run(): + controller = DecodeAdmissionController(capacity=1, max_queued_slots=1, timeout_seconds=1) + active = await controller.acquire(1) + waiting_task = asyncio.create_task(controller.acquire(1)) + 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_decode_admission_timeout_removes_waiter(): + async def run(): + controller = DecodeAdmissionController(capacity=1, max_queued_slots=1, timeout_seconds=0.01) + active = await controller.acquire(1) + + with pytest.raises(ServerBusyError, match="wait timed out"): + await controller.acquire(1) + + assert controller.queued_request_count == 0 + active.release() + + asyncio.run(run()) + + +def test_decode_admission_rejects_full_queue_and_oversized_gang(): + async def run(): + controller = DecodeAdmissionController(capacity=2, max_queued_slots=1, timeout_seconds=1) + active = await controller.acquire(2) + waiting_task = asyncio.create_task(controller.acquire(1)) + await asyncio.sleep(0) + + with pytest.raises(ServerBusyError, match="queue is full"): + await controller.acquire(1) + with pytest.raises(ServerBusyError, match="needs 3 Decode slots"): + await controller.acquire(3) + + waiting_task.cancel() + await asyncio.gather(waiting_task, return_exceptions=True) + active.release() + + asyncio.run(run()) + + +def test_decode_admission_cancellation_after_concurrent_grant_releases_lease(): + async def run(): + controller = DecodeAdmissionController(capacity=1, max_queued_slots=1, timeout_seconds=1) + active = await controller.acquire(1) + waiting_task = asyncio.create_task(controller.acquire(1)) + await asyncio.sleep(0) + + active.release() + waiting_task.cancel() + await asyncio.gather(waiting_task, return_exceptions=True) + + assert controller.active_slots == 0 + assert controller.queued_request_count == 0 + + asyncio.run(run()) + + +def test_decode_admission_gang_lease_can_transfer_to_independent_requests(): + async def run(): + controller = DecodeAdmissionController(capacity=3, max_queued_slots=0, timeout_seconds=1) + gang = await controller.acquire(3) + first, second, third = gang.split([1, 1, 1]) + + first.release() + assert controller.active_slots == 2 + second.release() + assert controller.active_slots == 1 + third.release() + assert controller.active_slots == 0 + + with pytest.raises(RuntimeError, match="already been released"): + gang.split([1, 1, 1]) + + asyncio.run(run()) + + +def test_decode_admission_handle_transfers_or_releases_lease_once(): + async def run(): + controller = DecodeAdmissionController(capacity=2, max_queued_slots=0, timeout_seconds=1) + transferred_handle = DecodeAdmissionLeaseHandle(await controller.acquire(1)) + lease = transferred_handle.take() + transferred_handle.release() + assert controller.active_slots == 1 + lease.release() + + abandoned_handle = DecodeAdmissionLeaseHandle(await controller.acquire(1)) + abandoned_handle.release() + abandoned_handle.release() + assert controller.active_slots == 0 + + asyncio.run(run()) diff --git a/unit_tests/server/httpserver/test_pd_generate_error.py b/unit_tests/server/httpserver/test_pd_generate_error.py index 3c7a4242a0..08ea61d1a7 100644 --- a/unit_tests/server/httpserver/test_pd_generate_error.py +++ b/unit_tests/server/httpserver/test_pd_generate_error.py @@ -7,13 +7,14 @@ import pytest from lightllm.server.core.objs import FinishStatus, SamplingParams -from lightllm.server.httpserver.pd_loop import _pd_process_generate +from lightllm.server.httpserver.decode_admission import DecodeAdmissionController +from lightllm.server.httpserver.pd_loop import _pd_process_generate, _reserve_decode_slots from lightllm.server.httpserver_for_pd_master.manager import ( HttpServerManagerForPDMaster, ReqStatus, ) from lightllm.server.pd_io_struct import ObjType -from lightllm.utils.error_utils import PDPrefillNodeStopGenToken +from lightllm.utils.error_utils import PDPrefillNodeStopGenToken, ServerBusyError class _FailingManager: @@ -32,6 +33,14 @@ async def generate(self, **_kwargs): yield +class _BusyManager: + args = SimpleNamespace(run_mode="decode") + + async def generate(self, **_kwargs): + raise ServerBusyError("decode queue is full") + yield + + class _SuccessfulManager: args = SimpleNamespace(run_mode="prefill") @@ -107,6 +116,28 @@ async def run(): asyncio.run(run()) +def test_pd_node_reports_server_busy_separately(): + async def run(): + sampling_params = SamplingParams() + sampling_params.group_request_id = 123 + websocket = AsyncMock() + + await _pd_process_generate( + manager=_BusyManager(), + prompt="prompt", + sampling_params=sampling_params, + multimodal_params=MagicMock(), + forwarding_queue=MagicMock(), + pd_upload_websocket=websocket, + pd_event=asyncio.Event(), + ) + + obj = pickle.loads(websocket.send.await_args.args[0]) + assert obj == (ObjType.PD_UPLOAD_SERVER_BUSY, 123, "decode queue is full") + + asyncio.run(run()) + + def test_pd_node_success_forwards_token_without_error_report(): async def run(): sampling_params = SamplingParams() @@ -204,6 +235,100 @@ async def run(): asyncio.run(run()) +def test_pd_decode_node_reserves_n_slots_and_hands_out_independent_leases(): + async def run(): + manager = SimpleNamespace( + decode_admission_controller=DecodeAdmissionController(capacity=3, max_queued_slots=0, timeout_seconds=1), + metric_client=MagicMock(), + ) + websocket = AsyncMock() + handles = {} + + await _reserve_decode_slots(manager, 100, (11, 22, 33), handles, websocket) + + assert manager.decode_admission_controller.active_slots == 3 + assert set(handles) == {11, 22, 33} + assert pickle.loads(websocket.send.await_args.args[0]) == (ObjType.PD_DECODE_SLOTS_RESERVED, 100) + + first = handles.pop(11).take() + first.release() + handles.pop(22).release() + handles.pop(33).release() + assert manager.decode_admission_controller.active_slots == 0 + + asyncio.run(run()) + + +def test_pd_decode_node_cancels_queued_atomic_reservation_without_leaking(): + async def run(): + controller = DecodeAdmissionController(capacity=2, max_queued_slots=2, timeout_seconds=10) + active = await controller.acquire(2) + manager = SimpleNamespace(decode_admission_controller=controller, metric_client=MagicMock()) + websocket = AsyncMock() + handles = {} + task = asyncio.create_task(_reserve_decode_slots(manager, 100, (11, 22), handles, websocket)) + await asyncio.sleep(0) + + task.cancel() + with suppress(asyncio.CancelledError): + await task + + assert controller.queued_slots == 0 + assert handles == {} + websocket.send.assert_not_awaited() + active.release() + assert controller.active_slots == 0 + + asyncio.run(run()) + + +def test_pd_master_routes_decode_reservation_ack_and_busy_response(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = SimpleNamespace(config_server_host=None, config_server_port=None) + manager.pd_manager = MagicMock() + manager.timer_log = AsyncMock() + manager.infos_queues = None + manager.req_id_to_out_inf = {} + manager.decode_reservation_statuses = {} + d_node = SimpleNamespace( + start_args={"pd_node_decode_admission_timeout": 1}, + websocket=SimpleNamespace(send_bytes=AsyncMock()), + ) + request = SimpleNamespace(is_disconnected=AsyncMock(return_value=False)) + + handle_task = asyncio.create_task(manager.handle_loop()) + try: + while manager.infos_queues is None: + await asyncio.sleep(0) + + accepted = asyncio.create_task(manager.reserve_decode_slots(d_node, 100, (11, 22), request)) + while 100 not in manager.decode_reservation_statuses: + await asyncio.sleep(0) + await manager.put_to_handle_queue((ObjType.PD_DECODE_SLOTS_RESERVED, 100)) + await asyncio.wait_for(accepted, timeout=1) + + rejected = asyncio.create_task(manager.reserve_decode_slots(d_node, 200, (33, 44), request)) + while 200 not in manager.decode_reservation_statuses: + await asyncio.sleep(0) + await manager.put_to_handle_queue((ObjType.PD_UPLOAD_SERVER_BUSY, 200, "decode queue is full")) + with pytest.raises(ServerBusyError, match="decode queue is full"): + await asyncio.wait_for(rejected, timeout=1) + + sent = [pickle.loads(call.args[0]) for call in d_node.websocket.send_bytes.await_args_list] + assert sent == [ + (ObjType.PD_RESERVE_DECODE_SLOTS, 100, (11, 22)), + (ObjType.PD_RESERVE_DECODE_SLOTS, 200, (33, 44)), + ] + assert manager.decode_reservation_statuses == {} + finally: + handle_task.cancel() + with suppress(asyncio.CancelledError): + await handle_task + + asyncio.run(run()) + + def test_pd_master_generate_error_marks_request_and_wakes_all_waiters(): async def run(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) @@ -211,6 +336,7 @@ async def run(): manager.pd_manager = MagicMock() manager.timer_log = AsyncMock() manager.infos_queues = None + manager.decode_reservation_statuses = {} p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) @@ -244,6 +370,33 @@ async def run(): asyncio.run(run()) +def test_pd_master_server_busy_wakes_waiters_and_preserves_429(): + async def run(): + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.args = SimpleNamespace(config_server_host=None) + manager.pd_manager = MagicMock() + manager.timer_log = AsyncMock() + manager.infos_queues = None + + req_status = ReqStatus(123, MagicMock(), MagicMock()) + manager.req_id_to_out_inf = {123: req_status} + handle_task = asyncio.create_task(manager.handle_loop()) + try: + while manager.infos_queues is None: + await asyncio.sleep(0) + await manager.put_to_handle_queue((ObjType.PD_UPLOAD_SERVER_BUSY, 123, "decode queue is full")) + await asyncio.wait_for(req_status.event.wait(), timeout=1) + + with pytest.raises(ServerBusyError, match="decode queue is full"): + req_status.raise_if_error() + finally: + handle_task.cancel() + with suppress(asyncio.CancelledError): + await handle_task + + asyncio.run(run()) + + @pytest.mark.parametrize("event_name", ["prefill_prompt_ids_event", "up_status_event"]) def test_pd_master_generate_error_wakes_resource_wait(event_name): async def run(): 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 da07759a99..6635cdfcd6 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -15,7 +15,6 @@ 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] @@ -40,7 +39,7 @@ def gen_id(): def _collect(mgr, sampling_params, monkeypatch, split): mgr._split_max_new_tokens = lambda *a, **k: list(split) - async def fake_wait(p_node, d_node, start_time, prompt, sp, multimodal_params, request): + async def fake_wait(p_node, d_node, start_time, prompt, sp, multimodal_params, request, decode_reservation=None): sub_req_id = sp.group_request_id hit = sp.max_new_tokens * 10 yield sub_req_id, "x", {"prompt_tokens": 100, "prompt_cache_len": hit}, FinishStatus() diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 879012c2ab..82e1d29fa4 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -6,6 +6,7 @@ import pytest from lightllm.server.core.objs import SamplingParams +from lightllm.server.httpserver.decode_admission import DecodeAdmissionController from lightllm.server.httpserver.manager import HttpServerManager from lightllm.server.pd_io_struct import NodeRole, ObjType from lightllm.utils.error_utils import PDPrefillNodeStopGenToken @@ -38,7 +39,9 @@ def _make_manager(mode: NodeRole): manager._register_running_request = AsyncMock() manager._unregister_running_request = AsyncMock() manager.metric_client = MagicMock() - manager.shm_req_manager = SimpleNamespace(async_alloc_req_index=AsyncMock(side_effect=RuntimeError("alloc failed"))) + manager.shm_req_manager = SimpleNamespace( + async_alloc_req_indexes=AsyncMock(side_effect=RuntimeError("alloc failed")) + ) return manager @@ -187,6 +190,68 @@ async def run(): asyncio.run(run()) +def test_decode_admission_lease_lives_until_request_resources_are_recycled(): + async def run(): + manager = _make_manager(NodeRole.D) + manager.decode_admission_controller = DecodeAdmissionController( + capacity=1, max_queued_slots=1, timeout_seconds=1 + ) + manager.args = SimpleNamespace(chunked_prefill_size=None) + manager.tokenizer = object() + manager.enable_multimodal = False + manager.transfer_to_next_module_or_node = AsyncMock() + manager._count_multimodal_tokens = MagicMock(return_value=(0, 0)) + + req_obj = MagicMock(request_id=123, index_in_shm_mem=7) + manager.shm_req_manager = SimpleNamespace( + async_alloc_req_indexes=AsyncMock(return_value=[7]), + async_get_req_obj_by_index=AsyncMock(return_value=req_obj), + async_put_back_req_obj=AsyncMock(), + async_release_req_index=AsyncMock(), + ) + + async def wait_to_token_package(*_args, **_kwargs): + yield 123, "token", {}, MagicMock() + + manager._wait_to_token_package = wait_to_token_package + generator = manager.generate( + prompt=[10, 11, 12], + sampling_params=_sampling_params(), + multimodal_params=_multimodal_params(), + request=None, + ) + + assert (await generator.__anext__())[1] == "token" + req_status = manager.req_id_to_out_inf[123] + assert manager.decode_admission_controller.active_slots == 1 + + await generator.aclose() + assert manager.decode_admission_controller.active_slots == 1 + + req_status.release_decode_admission() + assert manager.decode_admission_controller.active_slots == 0 + + asyncio.run(run()) + + +def test_multinode_tp_slave_does_not_register_as_independent_pd_node(): + async def run(): + manager = HttpServerManager.__new__(HttpServerManager) + manager.recycle_resource_loop = AsyncMock() + manager.loop_for_request = AsyncMock() + manager.is_multinode_tp_slave = True + manager.pd_mode = NodeRole.D + manager.zmq_recv_socket = SimpleNamespace(recv_pyobj=AsyncMock(side_effect=asyncio.CancelledError)) + + with patch("lightllm.server.httpserver.pd_loop.pd_handle_loop") as pd_handle_loop: + with pytest.raises(asyncio.CancelledError): + await manager.handle_loop() + + pd_handle_loop.assert_not_called() + + asyncio.run(run()) + + def test_running_request_helpers_are_atomic_and_refresh_timestamp_only_on_idle_to_running_transition(): async def run(): manager = HttpServerManager.__new__(HttpServerManager) diff --git a/unit_tests/server/test_pd_admission.py b/unit_tests/server/test_pd_admission.py deleted file mode 100644 index 802e293c0f..0000000000 --- a/unit_tests/server/test_pd_admission.py +++ /dev/null @@ -1,627 +0,0 @@ -import asyncio - -import pytest - -from lightllm.server.httpserver_for_pd_master.admission import ( - AdmissionPolicy, - AdmissionPriority, - AdmissionRequest, - PDAdmissionController, - SessionTracker, -) -from lightllm.utils.error_utils import ServerBusyError - - -def _request( - priority=AdmissionPriority.COLD, - session_key=None, - decode_slots=1, -): - return AdmissionRequest( - session_key=session_key, - priority=priority, - decode_slots=decode_slots, - ) - - -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_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) - 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_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)] - 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() - later = await later_task - assert multi_choice_task.done() is False - assert controller.active_slots == 3 - assert all(deficit >= 0 for deficit in controller._deficits.values()) - - active[1].release() - assert multi_choice_task.done() is False - - later.release() - multi_choice = await multi_choice_task - assert controller.active_slots == 3 - multi_choice.release() - active[2].release() - - asyncio.run(run()) - - -def test_blocked_gang_keeps_backfill_open_for_a_later_small_request(): - async def run(): - 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() - 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() - gang = await gang_task - gang.release() - - asyncio.run(run()) - - -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 - - 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 - - small = await controller.acquire(_request()) - assert controller.active_slots == 4 - assert controller.queued_slots == 4 - assert controller._backfilled_slots == 1 - - 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()) - - -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_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( - 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_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] - 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_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]) - - 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_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_requests_remain_fifo_within_the_same_priority_class(): - async def run(): - controller = PDAdmissionController( - lambda: 1, - policy=AdmissionPolicy(waiting_decode_waves=2), - ) - 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.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_cache_aware.py b/unit_tests/server/test_pd_cache_aware.py index 1e141110d5..6bdde78576 100644 --- a/unit_tests/server/test_pd_cache_aware.py +++ b/unit_tests/server/test_pd_cache_aware.py @@ -1,4 +1,3 @@ -import asyncio from types import SimpleNamespace import pytest @@ -74,70 +73,6 @@ 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_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) diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index 252d7b0c75..c748e47f7d 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -1,25 +1,14 @@ 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, - _build_pd_registration_info, - _update_pd_master_membership, -) -from lightllm.server.httpserver_for_pd_master.admission import ( - AdmissionPolicy, - AdmissionPriority, - PDAdmissionController, - SessionTracker, -) +from lightllm.server.httpserver.pd_loop import _build_pd_registration_info 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 +from lightllm.server.pd_io_struct import PD_DECODE_ADMISSION_CAPABILITY_KEY, NodeRole def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): @@ -197,455 +186,70 @@ 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_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() - 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_reports_only_actual_decode_capacity_changes(): +def test_pd_master_rejects_decode_node_without_authoritative_admission(): 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 - - 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 - assert ( - manager.update_node_load_info( + with pytest.raises(ValueError, match="upgrade Decode nodes before PD Masters"): + manager.register_pd( { - "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, - } + "node_id": 2, + "client_ip_port": "10.0.0.2:8000", + "mode": "decode", + "start_args": {"max_req_total_len": args.max_req_total_len}, + }, + websocket=object(), ) - is True - ) - node = manager.decode_nodes[0] - assert manager.get_decode_capacity() == 2 - assert node.run_status.total_token_usage_rate == 0.5 - # 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 + assert manager.decode_nodes == [] -def test_pd_registration_reads_capacity_from_legacy_safe_start_args(): +def test_pd_master_accepts_decode_node_with_authoritative_admission(): 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, - }, + "start_args": {"max_req_total_len": args.max_req_total_len, PD_DECODE_ADMISSION_CAPABILITY_KEY: True}, }, websocket=object(), ) - node = manager.decode_nodes[0] - assert node.capacity_share == 3 - assert node.capacity_epoch == 100 - assert manager.get_decode_capacity() == 3 + assert [node.client_ip_port for node in manager.decode_nodes] == ["10.0.0.2:8000"] -def test_pd_disconnect_accepts_transitional_top_level_capacity_fields(): - args = StartArgs() +def test_pd_master_allows_legacy_decode_node_only_when_admission_is_disabled(): + args = StartArgs(disable_pd_node_decode_admission=True) 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( - { - "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, + "start_args": {"max_req_total_len": args.max_req_total_len}, }, websocket=object(), ) - 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 - - -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 + assert [node.client_ip_port for node in manager.decode_nodes] == ["10.0.0.2:8000"] - 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" +def test_decode_registration_advertises_authoritative_admission(monkeypatch): + from lightllm.server.httpserver import pd_loop - manager._generate = fake_generate - generator = manager.generate("prompt", SimpleNamespace(n=1), None, None) - with pytest.raises(RuntimeError, match="metric enqueue failed"): - await generator.__anext__() + args = StartArgs(run_mode="decode", host="0.0.0.0") + manager = SimpleNamespace(args=args, host_ip="10.0.0.2", pd_mode=NodeRole.D, decode_admission_controller=object()) + monkeypatch.setattr(pd_loop, "get_shm_port_args", lambda: SimpleNamespace(port=8001)) - assert manager.admission_controller.active_slots == 0 - assert manager.running_request_count == 0 + registration = _build_pd_registration_info(manager) - asyncio.run(run()) + assert registration["start_args"][PD_DECODE_ADMISSION_CAPABILITY_KEY] is True + assert registration["start_args"]["host"] == "10.0.0.2" + assert args.host == "0.0.0.0" def test_prefill_registration_preserves_existing_inflight_prompt_chars(): @@ -726,80 +330,12 @@ def test_pd_master_inference_health_matches_normal_node_semantics(monkeypatch): assert HttpServerManagerForPDMaster.is_healthy(manager) is True -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 = [] - - async def fake_generate(prompt, *_args): - dispatched_prompts.append(prompt) - yield prompt - - manager._generate = fake_generate - first = manager.generate("first", None, None, None) - assert await first.__anext__() == "first" - - 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 - - await first.aclose() - assert await second_result == "second" - assert dispatched_prompts == ["first", "second"] - await second.aclose() - assert manager.admission_controller.active_slots == 0 - - asyncio.run(run()) - - -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), - ) - 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 - - 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(): 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(): @@ -814,7 +350,6 @@ 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