From 928d6bd10ece9e064de56ac6b1b1145a955e77ff Mon Sep 17 00:00:00 2001 From: sufubao Date: Sun, 9 Aug 2026 19:36:52 +0800 Subject: [PATCH 01/21] fix(pd): avoid first-token cache stats race --- lightllm/server/httpserver_for_pd_master/manager.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 96d1361e6..272b24522 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -428,6 +428,8 @@ async def fetch_pd_stream( ) first_token_gen = False + wait_prefill_first_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 + pending_decode_tokens = [] while True: await req_status.wait_to_ready() if await request.is_disconnected(): @@ -439,15 +441,23 @@ async def fetch_pd_stream( token_list = await req_status.pop_all_tokens() for sub_req_id, request_output, metadata, finish_status in token_list: output_index = metadata.get("count_output_tokens") + node_run_mode = metadata.get("node_mode") + if not first_token_gen and wait_prefill_first_token and node_run_mode != "prefill": + pending_decode_tokens.append((sub_req_id, request_output, metadata, finish_status)) + continue # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 if output_index == 1: if first_token_gen is False: first_token_gen = True - node_run_mode = metadata.pop("node_mode", None) + metadata.pop("node_mode", None) if node_run_mode == "prefill": if old_max_new_tokens != 1 and finish_status.is_finished_length(): finish_status = FinishStatus(FinishStatus.NO_FINISH) yield sub_req_id, request_output, metadata, finish_status + for pending_token in pending_decode_tokens: + if pending_token[2].get("count_output_tokens") != 1: + yield pending_token + pending_decode_tokens.clear() else: continue else: From e71f84a66448e22feb48799b6856839da2554ffa Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 10 Aug 2026 13:29:02 +0800 Subject: [PATCH 02/21] fix(pd): preserve prefill metadata and unblock decode --- .../httpserver_for_pd_master/manager.py | 17 ++- .../test_pd_master_cached_tokens.py | 109 ++++++++++++++++++ 2 files changed, 124 insertions(+), 2 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 272b24522..f894333bd 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -430,13 +430,21 @@ async def fetch_pd_stream( first_token_gen = False wait_prefill_first_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 pending_decode_tokens = [] + prefill_prompt_cache_len = None while True: - await req_status.wait_to_ready() + tokens_ready = await req_status.wait_to_ready() if await request.is_disconnected(): raise ClientDisconnected( group_request_id=group_request_id, reason="fetch_pd_stream decode period check network disconnected", ) + if not tokens_ready and any(token[3].is_finished() for token in pending_decode_tokens): + logger.warning( + f"group_request_id: {group_request_id} prefill first token missing; releasing decode output" + ) + for pending_token in pending_decode_tokens: + yield pending_token + return if await req_status.can_read(self.req_id_to_out_inf): token_list = await req_status.pop_all_tokens() for sub_req_id, request_output, metadata, finish_status in token_list: @@ -451,16 +459,20 @@ async def fetch_pd_stream( first_token_gen = True metadata.pop("node_mode", None) if node_run_mode == "prefill": + prefill_prompt_cache_len = metadata.get("prompt_cache_len", 0) if old_max_new_tokens != 1 and finish_status.is_finished_length(): finish_status = FinishStatus(FinishStatus.NO_FINISH) yield sub_req_id, request_output, metadata, finish_status for pending_token in pending_decode_tokens: if pending_token[2].get("count_output_tokens") != 1: + pending_token[2]["prompt_cache_len"] = prefill_prompt_cache_len yield pending_token pending_decode_tokens.clear() else: continue else: + if prefill_prompt_cache_len is not None: + metadata["prompt_cache_len"] = prefill_prompt_cache_len yield sub_req_id, request_output, metadata, finish_status return @@ -650,8 +662,9 @@ def __init__(self, req_id, p_node, d_node) -> None: async def wait_to_ready(self): try: await asyncio.wait_for(self.event.wait(), timeout=5) + return True except asyncio.TimeoutError: - pass + return False async def can_read(self, req_id_to_out_inf): async with self.lock: diff --git a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py index a3dc268a9..586fb53ba 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -1,13 +1,20 @@ import asyncio import copy +import pickle +from contextlib import aclosing from types import SimpleNamespace import pytest from lightllm.server.core.objs import FinishStatus, SamplingParams +from lightllm.server.httpserver_for_pd_master import manager as pd_master_manager from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster +def _ignore(*args): + pass + + def _make_manager(monkeypatch): monkeypatch.setattr( "lightllm.server.httpserver.manager.HttpServerManager._check_and_repair_length", @@ -85,3 +92,105 @@ def test_multi_block_keeps_first_block_hit(monkeypatch): cached = _collect(mgr, sp, monkeypatch, split=[3, 2]) assert cached[-1] == 30, cached assert mgr.recorded_cache_hit_rates == [pytest.approx(0.3)] + + +def _collect_fetch_pd_stream(monkeypatch, batches): + class FakeReqStatus: + def __init__(self): + self.up_status_event = asyncio.Event() + self.prefill_prompt_ids_event = asyncio.Event() + self.batches = list(batches) + + async def wait_to_ready(self): + return bool(self.batches) + + async def can_read(self, req_id_to_out_inf): + return bool(self.batches) + + async def pop_all_tokens(self): + return self.batches.pop(0) + + req_status = FakeReqStatus() + monkeypatch.setattr(pd_master_manager, "ReqStatus", lambda *args: req_status) + + manager = object.__new__(HttpServerManagerForPDMaster) + manager.args = SimpleNamespace(pd_node_id=0) + manager.req_id_to_out_inf = {} + + async def wait_for_stage(event, request, timeout, group_request_id, stage): + if stage == "prefill": + event.prompt_ids = [1, 2, 3, 4] + else: + decode_node_info = SimpleNamespace(ready_kv_len=0) + event.upkv_status = SimpleNamespace(pd_kv_trans_params=pickle.dumps(decode_node_info)) + + manager._wait_for_event_or_disconnect = wait_for_stage + websocket = SimpleNamespace(send_bytes=lambda data: asyncio.sleep(0)) + node = SimpleNamespace(websocket=websocket) + sampling_params = SimpleNamespace( + group_request_id=1, + max_new_tokens=3, + pd_master_node_id=SimpleNamespace(initialize=_ignore), + ) + request = SimpleNamespace(is_disconnected=lambda: asyncio.sleep(0, result=False)) + + async def run(): + results = [] + generator = manager.fetch_pd_stream(node, node, "prompt", sampling_params, None, request) + async with aclosing(generator): + async for result in generator: + results.append(result) + if result[3].is_finished(): + break + return results + + return asyncio.run(run()) + + +def test_prefill_cache_hit_is_copied_to_decode_tokens(monkeypatch): + batches = [ + [ + (1, "d1", {"count_output_tokens": 1, "node_mode": "decode", "prompt_cache_len": 1}, FinishStatus()), + (1, "d2", {"count_output_tokens": 2, "node_mode": "decode", "prompt_cache_len": 1}, FinishStatus()), + ], + [ + ( + 1, + "p1", + {"count_output_tokens": 1, "node_mode": "prefill", "prompt_cache_len": 8}, + FinishStatus(FinishStatus.FINISHED_LENGTH), + ) + ], + [ + ( + 1, + "d3", + {"count_output_tokens": 3, "node_mode": "decode", "prompt_cache_len": 1}, + FinishStatus(FinishStatus.FINISHED_STOP), + ) + ], + ] + + results = _collect_fetch_pd_stream(monkeypatch, batches) + + assert [result[1] for result in results] == ["p1", "d2", "d3"] + assert [result[2]["prompt_cache_len"] for result in results] == [8, 8, 8] + + +def test_finished_decode_is_released_when_prefill_token_is_missing(monkeypatch): + batches = [ + [ + (1, "d1", {"count_output_tokens": 1, "node_mode": "decode", "prompt_cache_len": 1}, FinishStatus()), + ( + 1, + "d2", + {"count_output_tokens": 2, "node_mode": "decode", "prompt_cache_len": 1}, + FinishStatus(FinishStatus.FINISHED_STOP), + ), + ] + ] + + results = _collect_fetch_pd_stream(monkeypatch, batches) + + assert [result[1] for result in results] == ["d1", "d2"] + assert results[-1][3].is_finished() From c912a635f5baf33e65d4a429daafacc0c6cde712 Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 10 Aug 2026 13:59:51 +0800 Subject: [PATCH 03/21] refactor(pd): simplify first-token fallback --- AGENTS.md | 149 ++++++++++++++++++ .../httpserver_for_pd_master/manager.py | 17 +- .../test_pd_master_cached_tokens.py | 109 ------------- 3 files changed, 156 insertions(+), 119 deletions(-) create mode 100644 AGENTS.md diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 000000000..2cb3d83e1 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,149 @@ +# AGENTS.md + +This file provides guidance to coding agents when working with code in this repository. + +## Working constitution (LightLLM only) + +These rules apply only to work in the LightLLM repository. Do not generalize them into global instructions or apply them to other repositories. + +### The user must fully understand the work + +Treat the user's understanding as a delivery requirement, not an optional explanation. Before implementation, establish and explain the observed behavior, relevant execution path, evidence, root cause or current uncertainty, and the proposed mechanism of change. Before handoff, explain what changed, why it fixes the problem, how it was verified, and any remaining uncertainty or tradeoff. Do not cross a material decision gate while the user indicates that they do not yet understand what is happening. + +### Implementation must follow Ponytail + +Before writing implementation code, apply the installed `ponytail` skill in its default full mode when it is available. If the skill is unavailable, follow its core policy directly: understand and trace the real flow first, then choose the smallest solution that actually works. Question unnecessary work (YAGNI), reuse existing project code before adding helpers, prefer the standard library or existing dependencies, fix the root cause at the shared path, minimize files and diff size, and avoid speculative abstractions, scaffolding, boilerplate, and cleverness. Do not add test code to a PR unless the user explicitly requests it; verify with existing checks instead. Never simplify away explicit requirements, trust-boundary validation, security, or error handling that prevents data loss. + +## Development commands + +### Install + +The documented development environment uses Python 3.10 (the package requires Python >=3.9.16): + +```bash +conda create -n lightllm python=3.10 -y +conda activate lightllm +pip install -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cu124 +python setup.py install +``` + +For an editable checkout, the Docker build uses: + +```bash +pip install -e . --no-cache-dir +``` + +Build the project image with: + +```bash +docker build -t lightllm-dev -f docker/Dockerfile . +``` + +### Run the server + +```bash +python -m lightllm.server.api_server --model_dir /path/to/model +``` + +The default HTTP endpoint is port 8000. A minimal request is: + +```bash +curl http://127.0.0.1:8000/generate \ + -H 'Content-Type: application/json' \ + -d '{"inputs":"What is AI?","parameters":{"max_new_tokens":17}}' +``` + +Most runtime behavior is selected through CLI flags defined in `lightllm/server/api_cli.py` and normalized into `StartArgs` in `lightllm/server/core/objs/start_args_type.py`. + +### Lint and format + +The pre-commit configuration is authoritative: Black 21.12b0 and flake8 6.1.0, both with a 120-column limit. + +```bash +pip install pre-commit +pre-commit install +pre-commit run --all-files +``` + +To check only files changed from a base revision: + +```bash +pre-commit run --files $(git diff --name-only HEAD) +``` + +Do not substitute a different system Black version; formatting can differ from the pinned hook. `format.py` is a legacy all-files autopep8 script, not the CI formatter. + +### Tests + +Pytest is not declared in `requirements.txt`; install it separately if needed. + +```bash +python -m pip install pytest +python -m pytest unit_tests -q +python -m pytest unit_tests/server/test_hypercorn_config.py -q +python -m pytest unit_tests/server/test_hypercorn_config.py::test_hypercorn_config_is_parsed -q +``` + +There is no repository-wide pytest configuration or canonical full-suite CI command. Many tests exercise CUDA kernels or distributed/model-specific paths and require suitable GPUs, libraries, and model artifacts; prefer the narrowest relevant test first. + +Every benchmark or performance experiment, including commands under `test/benchmark/` or `test/performance/`, must be recorded through the experiment ledger: + +```bash +exp -m "short purpose" +``` + +### Documentation + +```bash +python -m pip install -r docs/EN/requirements-docs.txt +make -C docs/EN html +``` + +The Chinese documentation has the analogous `docs/CN/Makefile`. + +## Architecture + +### Process topology and serving modes + +`lightllm/server/api_server.py` is the main entrypoint. It parses CLI arguments and dispatches by `run_mode`: normal serving, prefill/decode workers, PD master, config server, or visual-only. `lightllm/server/api_start.py` validates and derives runtime settings, then uses multiprocessing `spawn` to assemble the HTTP API, router, detokenizer, metrics/cache, and optional vision/audio processes. Shared-memory objects, ZeroMQ channels, and NCCL groups connect these components. + +The router is the boundary between request scheduling and GPU execution. `lightllm/server/router/manager.py` starts one model-inference subprocess per local global rank and communicates with each through Unix-socket RPyC. All ranks receive the same initialization arguments and must agree on profiled KV-cache capacity. + +### Backend and inference loop + +`lightllm/server/router/model_infer/model_rpc.py` selects a backend according to serving mode and features. Normal execution defaults to the chunked-prefill backend; PD, DP, reward, diverse, token-healing, and constrained-output modes select specialized backends under `lightllm/server/router/model_infer/mode_backend/`. + +`ModeBackend` initializes distributed groups, model configuration, the model instance, request/KV managers, and the prompt cache, then runs the scheduling/inference loops. The chunked-prefill implementation overlaps GPU forward/sampling, CPU state updates, post-processing, and scheduling for the next batch. Requests enter through shared memory on the node master and are broadcast to participating ranks. + +`ModelInput` and `ModelOutput` in `lightllm/server/router/model_infer/mode_backend/batch_objs.py` form the CPU/GPU batch boundary. Preprocessing in `generic_pre_process.py` performs prefix-cache lookup, request/KV allocation, and construction of prefill or decode inputs. + +### Model composition and execution + +Model implementations register themselves by Hugging Face `model_type` through decorators in `lightllm/models/registry.py`. Importing `lightllm.models` populates this registry. Model-specific classes are generally composition roots that select weight, layer-inference, and inference-state implementations before delegating lifecycle work to `TpPartBaseModel` in `lightllm/common/basemodel/basemodel.py`. + +`TpPartBaseModel` owns the ordered initialization pipeline: config repair and validation, quantization setup, weight objects, request/KV managers, inference layers, Hugging Face weight loading, attention backend setup, autotuning, and CUDA-graph capture. Forward execution dispatches prefill versus decode; eligible decode batches use captured CUDA graphs. Per-call sequence, cache, distributed, graph, and microbatch metadata lives in `InferStateInfo`. + +Weights are loaded from Hugging Face checkpoints, preferring safetensors, and tensor-parallel slices are applied during loading. Model-family code under `lightllm/models/` supplies architecture-specific weight mappings and kernels while common lifecycle and scheduling code remains shared. + +### Token-level KV and request management + +The core memory model is token-paged KV storage rather than per-request contiguous buffers. `MemoryManager` in `lightllm/common/kv_cache_mem_manager/` profiles rank-consistent capacity, owns the GPU KV tensor, and delegates free-slot bookkeeping to a pinned-CPU allocator. Specialized managers support normal, quantized, MLA/DSA, and other model-specific KV layouts. + +`ReqManager` separately allocates compact request IDs and maintains the GPU request-to-token-index table. Batch preprocessing allocates physical KV slots and writes those mappings before model forward. Finishing or pausing a request releases unshared slots; reusable prefixes may first be inserted into the dynamic prompt cache. + +`lightllm/server/router/dynamic_prompt/radix_cache.py` implements the ref-counted radix tree used for prefix reuse. It maps token segments to physical KV-slot segments and only evicts unreferenced leaves, returning their slots to `MemoryManager`. + +### Parallelism and communication + +The global world is divided into `dp` replicas; ranks within each replica form the tensor-parallel group. Rank topology and device setup are centralized in `lightllm/utils/dist_utils.py`. TP weight slicing and layer collectives live under `lightllm/common/basemodel/layer_weights/` and `lightllm/common/basemodel/layer_infer/`. + +`lightllm/distributed/communication_op.py` centralizes communication-group creation and collective dispatch. All-reduce prefers optimized implementations when enabled and falls back to NCCL. Optional groups support TP+SP overlap, cross-DP prefill balancing, and DeepEP expert parallelism for MoE models. + +### Where to make changes + +- API flags, launch modes, and process wiring: `lightllm/server/api_cli.py`, `api_server.py`, `api_start.py`. +- Scheduling, request lifecycle, and batch construction: `lightllm/server/router/` and `mode_backend/`. +- Shared model lifecycle, memory, and distributed execution: `lightllm/common/` and `lightllm/distributed/`. +- Architecture-specific weights and kernels: `lightllm/models//`. +- API behavior and protocol compatibility: `lightllm/server/httpserver/`. +- Focused unit tests: `unit_tests/`; benchmark and performance harnesses: `test/benchmark/` and `test/performance/`. diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index f894333bd..6cc5b6cc3 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -432,19 +432,12 @@ async def fetch_pd_stream( pending_decode_tokens = [] prefill_prompt_cache_len = None while True: - tokens_ready = await req_status.wait_to_ready() + await req_status.wait_to_ready() if await request.is_disconnected(): raise ClientDisconnected( group_request_id=group_request_id, reason="fetch_pd_stream decode period check network disconnected", ) - if not tokens_ready and any(token[3].is_finished() for token in pending_decode_tokens): - logger.warning( - f"group_request_id: {group_request_id} prefill first token missing; releasing decode output" - ) - for pending_token in pending_decode_tokens: - yield pending_token - return if await req_status.can_read(self.req_id_to_out_inf): token_list = await req_status.pop_all_tokens() for sub_req_id, request_output, metadata, finish_status in token_list: @@ -474,6 +467,11 @@ async def fetch_pd_stream( if prefill_prompt_cache_len is not None: metadata["prompt_cache_len"] = prefill_prompt_cache_len yield sub_req_id, request_output, metadata, finish_status + elif pending_decode_tokens and pending_decode_tokens[-1][3].is_finished(): + logger.warning(f"{group_request_id}: prefill token missing; releasing decode output") + for pending_token in pending_decode_tokens: + yield pending_token + return return @@ -662,9 +660,8 @@ def __init__(self, req_id, p_node, d_node) -> None: async def wait_to_ready(self): try: await asyncio.wait_for(self.event.wait(), timeout=5) - return True except asyncio.TimeoutError: - return False + pass async def can_read(self, req_id_to_out_inf): async with self.lock: 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 586fb53ba..a3dc268a9 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -1,20 +1,13 @@ import asyncio import copy -import pickle -from contextlib import aclosing from types import SimpleNamespace import pytest from lightllm.server.core.objs import FinishStatus, SamplingParams -from lightllm.server.httpserver_for_pd_master import manager as pd_master_manager from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster -def _ignore(*args): - pass - - def _make_manager(monkeypatch): monkeypatch.setattr( "lightllm.server.httpserver.manager.HttpServerManager._check_and_repair_length", @@ -92,105 +85,3 @@ def test_multi_block_keeps_first_block_hit(monkeypatch): cached = _collect(mgr, sp, monkeypatch, split=[3, 2]) assert cached[-1] == 30, cached assert mgr.recorded_cache_hit_rates == [pytest.approx(0.3)] - - -def _collect_fetch_pd_stream(monkeypatch, batches): - class FakeReqStatus: - def __init__(self): - self.up_status_event = asyncio.Event() - self.prefill_prompt_ids_event = asyncio.Event() - self.batches = list(batches) - - async def wait_to_ready(self): - return bool(self.batches) - - async def can_read(self, req_id_to_out_inf): - return bool(self.batches) - - async def pop_all_tokens(self): - return self.batches.pop(0) - - req_status = FakeReqStatus() - monkeypatch.setattr(pd_master_manager, "ReqStatus", lambda *args: req_status) - - manager = object.__new__(HttpServerManagerForPDMaster) - manager.args = SimpleNamespace(pd_node_id=0) - manager.req_id_to_out_inf = {} - - async def wait_for_stage(event, request, timeout, group_request_id, stage): - if stage == "prefill": - event.prompt_ids = [1, 2, 3, 4] - else: - decode_node_info = SimpleNamespace(ready_kv_len=0) - event.upkv_status = SimpleNamespace(pd_kv_trans_params=pickle.dumps(decode_node_info)) - - manager._wait_for_event_or_disconnect = wait_for_stage - websocket = SimpleNamespace(send_bytes=lambda data: asyncio.sleep(0)) - node = SimpleNamespace(websocket=websocket) - sampling_params = SimpleNamespace( - group_request_id=1, - max_new_tokens=3, - pd_master_node_id=SimpleNamespace(initialize=_ignore), - ) - request = SimpleNamespace(is_disconnected=lambda: asyncio.sleep(0, result=False)) - - async def run(): - results = [] - generator = manager.fetch_pd_stream(node, node, "prompt", sampling_params, None, request) - async with aclosing(generator): - async for result in generator: - results.append(result) - if result[3].is_finished(): - break - return results - - return asyncio.run(run()) - - -def test_prefill_cache_hit_is_copied_to_decode_tokens(monkeypatch): - batches = [ - [ - (1, "d1", {"count_output_tokens": 1, "node_mode": "decode", "prompt_cache_len": 1}, FinishStatus()), - (1, "d2", {"count_output_tokens": 2, "node_mode": "decode", "prompt_cache_len": 1}, FinishStatus()), - ], - [ - ( - 1, - "p1", - {"count_output_tokens": 1, "node_mode": "prefill", "prompt_cache_len": 8}, - FinishStatus(FinishStatus.FINISHED_LENGTH), - ) - ], - [ - ( - 1, - "d3", - {"count_output_tokens": 3, "node_mode": "decode", "prompt_cache_len": 1}, - FinishStatus(FinishStatus.FINISHED_STOP), - ) - ], - ] - - results = _collect_fetch_pd_stream(monkeypatch, batches) - - assert [result[1] for result in results] == ["p1", "d2", "d3"] - assert [result[2]["prompt_cache_len"] for result in results] == [8, 8, 8] - - -def test_finished_decode_is_released_when_prefill_token_is_missing(monkeypatch): - batches = [ - [ - (1, "d1", {"count_output_tokens": 1, "node_mode": "decode", "prompt_cache_len": 1}, FinishStatus()), - ( - 1, - "d2", - {"count_output_tokens": 2, "node_mode": "decode", "prompt_cache_len": 1}, - FinishStatus(FinishStatus.FINISHED_STOP), - ), - ] - ] - - results = _collect_fetch_pd_stream(monkeypatch, batches) - - assert [result[1] for result in results] == ["d1", "d2"] - assert results[-1][3].is_finished() From b05ff814cf0783673344be31c164ef4a0d9d8aaa Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 10 Aug 2026 14:04:18 +0800 Subject: [PATCH 04/21] chore: keep PD fix PR focused --- AGENTS.md | 149 ------------------------------------------------------ 1 file changed, 149 deletions(-) delete mode 100644 AGENTS.md diff --git a/AGENTS.md b/AGENTS.md deleted file mode 100644 index 2cb3d83e1..000000000 --- a/AGENTS.md +++ /dev/null @@ -1,149 +0,0 @@ -# AGENTS.md - -This file provides guidance to coding agents when working with code in this repository. - -## Working constitution (LightLLM only) - -These rules apply only to work in the LightLLM repository. Do not generalize them into global instructions or apply them to other repositories. - -### The user must fully understand the work - -Treat the user's understanding as a delivery requirement, not an optional explanation. Before implementation, establish and explain the observed behavior, relevant execution path, evidence, root cause or current uncertainty, and the proposed mechanism of change. Before handoff, explain what changed, why it fixes the problem, how it was verified, and any remaining uncertainty or tradeoff. Do not cross a material decision gate while the user indicates that they do not yet understand what is happening. - -### Implementation must follow Ponytail - -Before writing implementation code, apply the installed `ponytail` skill in its default full mode when it is available. If the skill is unavailable, follow its core policy directly: understand and trace the real flow first, then choose the smallest solution that actually works. Question unnecessary work (YAGNI), reuse existing project code before adding helpers, prefer the standard library or existing dependencies, fix the root cause at the shared path, minimize files and diff size, and avoid speculative abstractions, scaffolding, boilerplate, and cleverness. Do not add test code to a PR unless the user explicitly requests it; verify with existing checks instead. Never simplify away explicit requirements, trust-boundary validation, security, or error handling that prevents data loss. - -## Development commands - -### Install - -The documented development environment uses Python 3.10 (the package requires Python >=3.9.16): - -```bash -conda create -n lightllm python=3.10 -y -conda activate lightllm -pip install -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cu124 -python setup.py install -``` - -For an editable checkout, the Docker build uses: - -```bash -pip install -e . --no-cache-dir -``` - -Build the project image with: - -```bash -docker build -t lightllm-dev -f docker/Dockerfile . -``` - -### Run the server - -```bash -python -m lightllm.server.api_server --model_dir /path/to/model -``` - -The default HTTP endpoint is port 8000. A minimal request is: - -```bash -curl http://127.0.0.1:8000/generate \ - -H 'Content-Type: application/json' \ - -d '{"inputs":"What is AI?","parameters":{"max_new_tokens":17}}' -``` - -Most runtime behavior is selected through CLI flags defined in `lightllm/server/api_cli.py` and normalized into `StartArgs` in `lightllm/server/core/objs/start_args_type.py`. - -### Lint and format - -The pre-commit configuration is authoritative: Black 21.12b0 and flake8 6.1.0, both with a 120-column limit. - -```bash -pip install pre-commit -pre-commit install -pre-commit run --all-files -``` - -To check only files changed from a base revision: - -```bash -pre-commit run --files $(git diff --name-only HEAD) -``` - -Do not substitute a different system Black version; formatting can differ from the pinned hook. `format.py` is a legacy all-files autopep8 script, not the CI formatter. - -### Tests - -Pytest is not declared in `requirements.txt`; install it separately if needed. - -```bash -python -m pip install pytest -python -m pytest unit_tests -q -python -m pytest unit_tests/server/test_hypercorn_config.py -q -python -m pytest unit_tests/server/test_hypercorn_config.py::test_hypercorn_config_is_parsed -q -``` - -There is no repository-wide pytest configuration or canonical full-suite CI command. Many tests exercise CUDA kernels or distributed/model-specific paths and require suitable GPUs, libraries, and model artifacts; prefer the narrowest relevant test first. - -Every benchmark or performance experiment, including commands under `test/benchmark/` or `test/performance/`, must be recorded through the experiment ledger: - -```bash -exp -m "short purpose" -``` - -### Documentation - -```bash -python -m pip install -r docs/EN/requirements-docs.txt -make -C docs/EN html -``` - -The Chinese documentation has the analogous `docs/CN/Makefile`. - -## Architecture - -### Process topology and serving modes - -`lightllm/server/api_server.py` is the main entrypoint. It parses CLI arguments and dispatches by `run_mode`: normal serving, prefill/decode workers, PD master, config server, or visual-only. `lightllm/server/api_start.py` validates and derives runtime settings, then uses multiprocessing `spawn` to assemble the HTTP API, router, detokenizer, metrics/cache, and optional vision/audio processes. Shared-memory objects, ZeroMQ channels, and NCCL groups connect these components. - -The router is the boundary between request scheduling and GPU execution. `lightllm/server/router/manager.py` starts one model-inference subprocess per local global rank and communicates with each through Unix-socket RPyC. All ranks receive the same initialization arguments and must agree on profiled KV-cache capacity. - -### Backend and inference loop - -`lightllm/server/router/model_infer/model_rpc.py` selects a backend according to serving mode and features. Normal execution defaults to the chunked-prefill backend; PD, DP, reward, diverse, token-healing, and constrained-output modes select specialized backends under `lightllm/server/router/model_infer/mode_backend/`. - -`ModeBackend` initializes distributed groups, model configuration, the model instance, request/KV managers, and the prompt cache, then runs the scheduling/inference loops. The chunked-prefill implementation overlaps GPU forward/sampling, CPU state updates, post-processing, and scheduling for the next batch. Requests enter through shared memory on the node master and are broadcast to participating ranks. - -`ModelInput` and `ModelOutput` in `lightllm/server/router/model_infer/mode_backend/batch_objs.py` form the CPU/GPU batch boundary. Preprocessing in `generic_pre_process.py` performs prefix-cache lookup, request/KV allocation, and construction of prefill or decode inputs. - -### Model composition and execution - -Model implementations register themselves by Hugging Face `model_type` through decorators in `lightllm/models/registry.py`. Importing `lightllm.models` populates this registry. Model-specific classes are generally composition roots that select weight, layer-inference, and inference-state implementations before delegating lifecycle work to `TpPartBaseModel` in `lightllm/common/basemodel/basemodel.py`. - -`TpPartBaseModel` owns the ordered initialization pipeline: config repair and validation, quantization setup, weight objects, request/KV managers, inference layers, Hugging Face weight loading, attention backend setup, autotuning, and CUDA-graph capture. Forward execution dispatches prefill versus decode; eligible decode batches use captured CUDA graphs. Per-call sequence, cache, distributed, graph, and microbatch metadata lives in `InferStateInfo`. - -Weights are loaded from Hugging Face checkpoints, preferring safetensors, and tensor-parallel slices are applied during loading. Model-family code under `lightllm/models/` supplies architecture-specific weight mappings and kernels while common lifecycle and scheduling code remains shared. - -### Token-level KV and request management - -The core memory model is token-paged KV storage rather than per-request contiguous buffers. `MemoryManager` in `lightllm/common/kv_cache_mem_manager/` profiles rank-consistent capacity, owns the GPU KV tensor, and delegates free-slot bookkeeping to a pinned-CPU allocator. Specialized managers support normal, quantized, MLA/DSA, and other model-specific KV layouts. - -`ReqManager` separately allocates compact request IDs and maintains the GPU request-to-token-index table. Batch preprocessing allocates physical KV slots and writes those mappings before model forward. Finishing or pausing a request releases unshared slots; reusable prefixes may first be inserted into the dynamic prompt cache. - -`lightllm/server/router/dynamic_prompt/radix_cache.py` implements the ref-counted radix tree used for prefix reuse. It maps token segments to physical KV-slot segments and only evicts unreferenced leaves, returning their slots to `MemoryManager`. - -### Parallelism and communication - -The global world is divided into `dp` replicas; ranks within each replica form the tensor-parallel group. Rank topology and device setup are centralized in `lightllm/utils/dist_utils.py`. TP weight slicing and layer collectives live under `lightllm/common/basemodel/layer_weights/` and `lightllm/common/basemodel/layer_infer/`. - -`lightllm/distributed/communication_op.py` centralizes communication-group creation and collective dispatch. All-reduce prefers optimized implementations when enabled and falls back to NCCL. Optional groups support TP+SP overlap, cross-DP prefill balancing, and DeepEP expert parallelism for MoE models. - -### Where to make changes - -- API flags, launch modes, and process wiring: `lightllm/server/api_cli.py`, `api_server.py`, `api_start.py`. -- Scheduling, request lifecycle, and batch construction: `lightllm/server/router/` and `mode_backend/`. -- Shared model lifecycle, memory, and distributed execution: `lightllm/common/` and `lightllm/distributed/`. -- Architecture-specific weights and kernels: `lightllm/models//`. -- API behavior and protocol compatibility: `lightllm/server/httpserver/`. -- Focused unit tests: `unit_tests/`; benchmark and performance harnesses: `test/benchmark/` and `test/performance/`. From 00f2698256210d8a3f1b0825c6f9f2011cd3535c Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 10 Aug 2026 14:11:35 +0800 Subject: [PATCH 05/21] refactor(pd): consume node mode once --- lightllm/server/httpserver_for_pd_master/manager.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 6cc5b6cc3..89d53a0f4 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -442,7 +442,7 @@ async def fetch_pd_stream( token_list = await req_status.pop_all_tokens() for sub_req_id, request_output, metadata, finish_status in token_list: output_index = metadata.get("count_output_tokens") - node_run_mode = metadata.get("node_mode") + node_run_mode = metadata.pop("node_mode", None) if not first_token_gen and wait_prefill_first_token and node_run_mode != "prefill": pending_decode_tokens.append((sub_req_id, request_output, metadata, finish_status)) continue @@ -450,7 +450,6 @@ async def fetch_pd_stream( if output_index == 1: if first_token_gen is False: first_token_gen = True - metadata.pop("node_mode", None) if node_run_mode == "prefill": prefill_prompt_cache_len = metadata.get("prompt_cache_len", 0) if old_max_new_tokens != 1 and finish_status.is_finished_length(): From e4604cb43aaafa0f247d6b4df03da7a6d3e199da Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 10 Aug 2026 14:52:11 +0800 Subject: [PATCH 06/21] refactor(pd): clarify first-token buffering --- .../httpserver_for_pd_master/manager.py | 73 ++++++++++--------- 1 file changed, 40 insertions(+), 33 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 89d53a0f4..adb79f05a 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -427,9 +427,9 @@ async def fetch_pd_stream( pickle.dumps((ObjType.PD_REQ_DECODE_NODE_INFO, group_request_id, decode_node_info)) ) - first_token_gen = False - wait_prefill_first_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 - pending_decode_tokens = [] + first_token_emitted = False + waiting_for_prefill_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 + token_list = [] prefill_prompt_cache_len = None while True: await req_status.wait_to_ready() @@ -439,38 +439,45 @@ async def fetch_pd_stream( reason="fetch_pd_stream decode period check network disconnected", ) if await req_status.can_read(self.req_id_to_out_inf): - token_list = await req_status.pop_all_tokens() - for sub_req_id, request_output, metadata, finish_status in token_list: - output_index = metadata.get("count_output_tokens") - node_run_mode = metadata.pop("node_mode", None) - if not first_token_gen and wait_prefill_first_token and node_run_mode != "prefill": - pending_decode_tokens.append((sub_req_id, request_output, metadata, finish_status)) - continue - # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 - if output_index == 1: - if first_token_gen is False: - first_token_gen = True - if node_run_mode == "prefill": - prefill_prompt_cache_len = metadata.get("prompt_cache_len", 0) - if old_max_new_tokens != 1 and finish_status.is_finished_length(): - finish_status = FinishStatus(FinishStatus.NO_FINISH) - yield sub_req_id, request_output, metadata, finish_status - for pending_token in pending_decode_tokens: - if pending_token[2].get("count_output_tokens") != 1: - pending_token[2]["prompt_cache_len"] = prefill_prompt_cache_len - yield pending_token - pending_decode_tokens.clear() - else: - continue - else: - if prefill_prompt_cache_len is not None: - metadata["prompt_cache_len"] = prefill_prompt_cache_len - yield sub_req_id, request_output, metadata, finish_status - elif pending_decode_tokens and pending_decode_tokens[-1][3].is_finished(): + token_list.extend(await req_status.pop_all_tokens()) + elif token_list and token_list[-1][3].is_finished(): + # D 已完成且一个轮询周期内仍未收到 P 首 token,用完整 D 输出兜底。 logger.warning(f"{group_request_id}: prefill token missing; releasing decode output") - for pending_token in pending_decode_tokens: - yield pending_token + for token in token_list: + token[2].pop("node_mode", None) + yield token return + else: + continue + + # 需要 prefill 时先累计 D 输出,直到拿到带权威缓存统计的 P 首 token。 + if waiting_for_prefill_token: + prefill_index = next( + (index for index, token in enumerate(token_list) if token[2].get("node_mode") == "prefill"), + None, + ) + if prefill_index is None: + continue + # P 首 token 必须先输出,后面的统一逻辑会丢弃重复的 D 首 token。 + prefill_token = token_list.pop(prefill_index) + token_list.insert(0, prefill_token) + waiting_for_prefill_token = False + + for sub_req_id, request_output, metadata, finish_status in token_list: + output_index = metadata.get("count_output_tokens") + node_run_mode = metadata.pop("node_mode", None) + if output_index == 1: + if first_token_emitted: + continue + first_token_emitted = True + if node_run_mode == "prefill": + prefill_prompt_cache_len = metadata.get("prompt_cache_len", 0) + if old_max_new_tokens != 1 and finish_status.is_finished_length(): + finish_status = FinishStatus(FinishStatus.NO_FINISH) + if prefill_prompt_cache_len is not None: + metadata["prompt_cache_len"] = prefill_prompt_cache_len + yield sub_req_id, request_output, metadata, finish_status + token_list.clear() return From 98d8626a7c93c362fab03d448d5cd5dc5530a5cd Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 10 Aug 2026 15:08:52 +0800 Subject: [PATCH 07/21] refactor(pd): make prefill lookup explicit --- .../httpserver_for_pd_master/manager.py | 25 +++++++++++-------- 1 file changed, 15 insertions(+), 10 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index adb79f05a..94c98730d 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -440,22 +440,27 @@ async def fetch_pd_stream( ) if await req_status.can_read(self.req_id_to_out_inf): token_list.extend(await req_status.pop_all_tokens()) - elif token_list and token_list[-1][3].is_finished(): + else: + if not token_list: + continue + _, _, _, last_finish_status = token_list[-1] + if not last_finish_status.is_finished(): + continue # D 已完成且一个轮询周期内仍未收到 P 首 token,用完整 D 输出兜底。 logger.warning(f"{group_request_id}: prefill token missing; releasing decode output") - for token in token_list: - token[2].pop("node_mode", None) - yield token + for sub_req_id, request_output, metadata, finish_status in token_list: + metadata.pop("node_mode", None) + yield sub_req_id, request_output, metadata, finish_status return - else: - continue # 需要 prefill 时先累计 D 输出,直到拿到带权威缓存统计的 P 首 token。 if waiting_for_prefill_token: - prefill_index = next( - (index for index, token in enumerate(token_list) if token[2].get("node_mode") == "prefill"), - None, - ) + prefill_index = None + for index, token in enumerate(token_list): + _, _, metadata, _ = token + if metadata.get("node_mode") == "prefill": + prefill_index = index + break if prefill_index is None: continue # P 首 token 必须先输出,后面的统一逻辑会丢弃重复的 D 首 token。 From 29ce20c13dee20b6d25e38796254d9615d40771f Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 10 Aug 2026 15:17:05 +0800 Subject: [PATCH 08/21] refactor(pd): avoid rescanning buffered tokens --- .../httpserver_for_pd_master/manager.py | 20 ++++++++++--------- 1 file changed, 11 insertions(+), 9 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 94c98730d..2288b7253 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -430,6 +430,7 @@ async def fetch_pd_stream( first_token_emitted = False waiting_for_prefill_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 token_list = [] + prefill_token = None prefill_prompt_cache_len = None while True: await req_status.wait_to_ready() @@ -439,7 +440,7 @@ async def fetch_pd_stream( reason="fetch_pd_stream decode period check network disconnected", ) if await req_status.can_read(self.req_id_to_out_inf): - token_list.extend(await req_status.pop_all_tokens()) + new_tokens = await req_status.pop_all_tokens() else: if not token_list: continue @@ -453,20 +454,21 @@ async def fetch_pd_stream( yield sub_req_id, request_output, metadata, finish_status return - # 需要 prefill 时先累计 D 输出,直到拿到带权威缓存统计的 P 首 token。 + # 只检查本轮新 token,避免等待 P 时反复扫描已缓存的 D 输出。 if waiting_for_prefill_token: - prefill_index = None - for index, token in enumerate(token_list): + for token in new_tokens: _, _, metadata, _ = token - if metadata.get("node_mode") == "prefill": - prefill_index = index - break - if prefill_index is None: + if prefill_token is None and metadata.get("node_mode") == "prefill": + prefill_token = token + else: + token_list.append(token) + if prefill_token is None: continue # P 首 token 必须先输出,后面的统一逻辑会丢弃重复的 D 首 token。 - prefill_token = token_list.pop(prefill_index) token_list.insert(0, prefill_token) waiting_for_prefill_token = False + else: + token_list.extend(new_tokens) for sub_req_id, request_output, metadata, finish_status in token_list: output_index = metadata.get("count_output_tokens") From c979dd0331c30b55ce37b55567f1d29483cb1f39 Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 10 Aug 2026 19:20:45 +0800 Subject: [PATCH 09/21] refactor(pd): clarify first-token stream handling --- .../httpserver_for_pd_master/manager.py | 104 +++++++++--------- 1 file changed, 54 insertions(+), 50 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 2288b7253..10872e5fe 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -427,50 +427,63 @@ async def fetch_pd_stream( pickle.dumps((ObjType.PD_REQ_DECODE_NODE_INFO, group_request_id, decode_node_info)) ) - first_token_emitted = False - waiting_for_prefill_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 - token_list = [] - prefill_token = None - prefill_prompt_cache_len = None - while True: - await req_status.wait_to_ready() - if await request.is_disconnected(): - raise ClientDisconnected( - group_request_id=group_request_id, - reason="fetch_pd_stream decode period check network disconnected", - ) - if await req_status.can_read(self.req_id_to_out_inf): - new_tokens = await req_status.pop_all_tokens() - else: - if not token_list: + needs_prefill_first_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 + buffered_decode_tokens = [] + initial_tokens = [] + + # 阶段 1:缓存先到的 D 输出,直到收到携带缓存统计的 P 首 token。 + if needs_prefill_first_token: + while True: + new_tokens = await req_status.drain_tokens(self.req_id_to_out_inf) + if await request.is_disconnected(): + raise ClientDisconnected( + group_request_id=group_request_id, + reason="fetch_pd_stream decode period check network disconnected", + ) + + if new_tokens: + prefill_token = None + for token in new_tokens: + _, _, metadata, _ = token + if prefill_token is None and metadata.get("node_mode") == "prefill": + prefill_token = token + else: + buffered_decode_tokens.append(token) + if prefill_token is None: + continue + # P 首 token 必须先输出,后面的统一逻辑会丢弃重复的 D 首 token。 + initial_tokens = [prefill_token, *buffered_decode_tokens] + buffered_decode_tokens.clear() + break + + if not buffered_decode_tokens: continue - _, _, _, last_finish_status = token_list[-1] + _, _, _, last_finish_status = buffered_decode_tokens[-1] if not last_finish_status.is_finished(): continue # D 已完成且一个轮询周期内仍未收到 P 首 token,用完整 D 输出兜底。 logger.warning(f"{group_request_id}: prefill token missing; releasing decode output") - for sub_req_id, request_output, metadata, finish_status in token_list: + for sub_req_id, request_output, metadata, finish_status in buffered_decode_tokens: metadata.pop("node_mode", None) yield sub_req_id, request_output, metadata, finish_status return - # 只检查本轮新 token,避免等待 P 时反复扫描已缓存的 D 输出。 - if waiting_for_prefill_token: - for token in new_tokens: - _, _, metadata, _ = token - if prefill_token is None and metadata.get("node_mode") == "prefill": - prefill_token = token - else: - token_list.append(token) - if prefill_token is None: + # 阶段 2:输出合并后的 token 流,并将 P 的缓存统计传递给后续 D token。 + first_token_emitted = False + prompt_cache_len_from_prefill = None + token_batch = initial_tokens + while True: + if not token_batch: + token_batch = await req_status.drain_tokens(self.req_id_to_out_inf) + if await request.is_disconnected(): + raise ClientDisconnected( + group_request_id=group_request_id, + reason="fetch_pd_stream decode period check network disconnected", + ) + if not token_batch: continue - # P 首 token 必须先输出,后面的统一逻辑会丢弃重复的 D 首 token。 - token_list.insert(0, prefill_token) - waiting_for_prefill_token = False - else: - token_list.extend(new_tokens) - for sub_req_id, request_output, metadata, finish_status in token_list: + for sub_req_id, request_output, metadata, finish_status in token_batch: output_index = metadata.get("count_output_tokens") node_run_mode = metadata.pop("node_mode", None) if output_index == 1: @@ -478,13 +491,13 @@ async def fetch_pd_stream( continue first_token_emitted = True if node_run_mode == "prefill": - prefill_prompt_cache_len = metadata.get("prompt_cache_len", 0) + prompt_cache_len_from_prefill = metadata.get("prompt_cache_len", 0) if old_max_new_tokens != 1 and finish_status.is_finished_length(): finish_status = FinishStatus(FinishStatus.NO_FINISH) - if prefill_prompt_cache_len is not None: - metadata["prompt_cache_len"] = prefill_prompt_cache_len + if prompt_cache_len_from_prefill is not None: + metadata["prompt_cache_len"] = prompt_cache_len_from_prefill yield sub_req_id, request_output, metadata, finish_status - token_list.clear() + token_batch = [] return @@ -670,26 +683,17 @@ 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 - async def wait_to_ready(self): + async def drain_tokens(self, req_id_to_out_inf): try: await asyncio.wait_for(self.event.wait(), timeout=5) except asyncio.TimeoutError: pass - - async def can_read(self, req_id_to_out_inf): async with self.lock: self.event.clear() assert self.req_id in req_id_to_out_inf, f"error state req_id {self.req_id}" - if len(self.out_token_info_list) == 0: - return False - else: - return True - - async def pop_all_tokens(self): - async with self.lock: - ans = self.out_token_info_list.copy() - self.out_token_info_list.clear() - return ans + tokens = self.out_token_info_list + self.out_token_info_list = [] + return tokens class PDManager: From 750a02edb98716ce7d254416168911d87271ae4d Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 10 Aug 2026 19:41:05 +0800 Subject: [PATCH 10/21] refactor(pd): streamline first-token buffering --- .../httpserver_for_pd_master/manager.py | 69 +++++++------------ 1 file changed, 23 insertions(+), 46 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 10872e5fe..fb8c950e5 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -427,63 +427,41 @@ async def fetch_pd_stream( pickle.dumps((ObjType.PD_REQ_DECODE_NODE_INFO, group_request_id, decode_node_info)) ) + first_token_emitted = False needs_prefill_first_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 buffered_decode_tokens = [] - initial_tokens = [] - - # 阶段 1:缓存先到的 D 输出,直到收到携带缓存统计的 P 首 token。 - if needs_prefill_first_token: - while True: - new_tokens = await req_status.drain_tokens(self.req_id_to_out_inf) - if await request.is_disconnected(): - raise ClientDisconnected( - group_request_id=group_request_id, - reason="fetch_pd_stream decode period check network disconnected", - ) - - if new_tokens: - prefill_token = None - for token in new_tokens: - _, _, metadata, _ = token - if prefill_token is None and metadata.get("node_mode") == "prefill": - prefill_token = token - else: - buffered_decode_tokens.append(token) - if prefill_token is None: - continue - # P 首 token 必须先输出,后面的统一逻辑会丢弃重复的 D 首 token。 - initial_tokens = [prefill_token, *buffered_decode_tokens] - buffered_decode_tokens.clear() - break + prompt_cache_len_from_prefill = None + while True: + new_tokens = await req_status.drain_tokens(self.req_id_to_out_inf) + if await request.is_disconnected(): + raise ClientDisconnected( + group_request_id=group_request_id, + reason="fetch_pd_stream decode period check network disconnected", + ) - if not buffered_decode_tokens: + if not new_tokens: + if not buffered_decode_tokens or not buffered_decode_tokens[-1][3].is_finished(): continue - _, _, _, last_finish_status = buffered_decode_tokens[-1] - if not last_finish_status.is_finished(): - continue - # D 已完成且一个轮询周期内仍未收到 P 首 token,用完整 D 输出兜底。 logger.warning(f"{group_request_id}: prefill token missing; releasing decode output") for sub_req_id, request_output, metadata, finish_status in buffered_decode_tokens: metadata.pop("node_mode", None) yield sub_req_id, request_output, metadata, finish_status return - # 阶段 2:输出合并后的 token 流,并将 P 的缓存统计传递给后续 D token。 - first_token_emitted = False - prompt_cache_len_from_prefill = None - token_batch = initial_tokens - while True: - if not token_batch: - token_batch = await req_status.drain_tokens(self.req_id_to_out_inf) - if await request.is_disconnected(): - raise ClientDisconnected( - group_request_id=group_request_id, - reason="fetch_pd_stream decode period check network disconnected", - ) - if not token_batch: + # P 首 token 携带缓存统计,需先于已缓冲的 D 输出返回。 + if needs_prefill_first_token: + prefill_token = next( + (token for token in new_tokens if token[2].get("node_mode") == "prefill"), None + ) + if prefill_token is None: + buffered_decode_tokens.extend(new_tokens) continue + new_tokens.remove(prefill_token) + new_tokens = [prefill_token, *buffered_decode_tokens, *new_tokens] + buffered_decode_tokens.clear() + needs_prefill_first_token = False - for sub_req_id, request_output, metadata, finish_status in token_batch: + for sub_req_id, request_output, metadata, finish_status in new_tokens: output_index = metadata.get("count_output_tokens") node_run_mode = metadata.pop("node_mode", None) if output_index == 1: @@ -497,7 +475,6 @@ async def fetch_pd_stream( if prompt_cache_len_from_prefill is not None: metadata["prompt_cache_len"] = prompt_cache_len_from_prefill yield sub_req_id, request_output, metadata, finish_status - token_batch = [] return From 563e4ddf043c8ecea429b98c8f6320566fbbb91b Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 10 Aug 2026 20:40:05 +0800 Subject: [PATCH 11/21] refactor(pd): reuse async queue for output tokens --- lightllm/server/httpserver/async_queue.py | 8 ++--- .../httpserver_for_pd_master/manager.py | 29 +++++-------------- 2 files changed, 11 insertions(+), 26 deletions(-) diff --git a/lightllm/server/httpserver/async_queue.py b/lightllm/server/httpserver/async_queue.py index a9f0c9068..866d75726 100644 --- a/lightllm/server/httpserver/async_queue.py +++ b/lightllm/server/httpserver/async_queue.py @@ -7,9 +7,9 @@ def __init__(self): self.event = asyncio.Event() self.lock = asyncio.Lock() - async def wait_to_ready(self): + async def wait_to_ready(self, timeout=3): try: - await asyncio.wait_for(self.event.wait(), timeout=3) + await asyncio.wait_for(self.event.wait(), timeout=timeout) except asyncio.TimeoutError: pass @@ -26,7 +26,7 @@ async def put(self, obj): self.event.set() return - async def wait_to_get_all_data(self): - await self.wait_to_ready() + async def wait_to_get_all_data(self, timeout=3): + await self.wait_to_ready(timeout=timeout) handle_list = await self.get_all_data() return handle_list diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index fb8c950e5..624971dc0 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -432,13 +432,15 @@ async def fetch_pd_stream( buffered_decode_tokens = [] prompt_cache_len_from_prefill = None while True: - new_tokens = await req_status.drain_tokens(self.req_id_to_out_inf) + new_tokens = await req_status.out_tokens.wait_to_get_all_data(timeout=5) + assert group_request_id in self.req_id_to_out_inf, f"error state req_id {group_request_id}" if await request.is_disconnected(): raise ClientDisconnected( group_request_id=group_request_id, reason="fetch_pd_stream decode period check network disconnected", ) + # 本轮无新 token 且 D 已完成时,释放缓冲避免因 P 首 token 缺失而挂起。 if not new_tokens: if not buffered_decode_tokens or not buffered_decode_tokens[-1][3].is_finished(): continue @@ -618,18 +620,15 @@ async def handle_loop(self): group_req_id = convert_sub_id_to_group_id(sub_req_id) try: req_status: ReqStatus = self.req_id_to_out_inf[group_req_id] - async with req_status.lock: - req_status.out_token_info_list.append((sub_req_id, text, metadata, finish_status)) - req_status.event.set() + await req_status.out_tokens.put((sub_req_id, text, metadata, finish_status)) except: pass elif obj[0] == ObjType.PD_UPLOAD_PREFILL_PROMPT_IDS: _, group_req_id, prompt_ids = obj try: req_status: ReqStatus = self.req_id_to_out_inf[group_req_id] - async with req_status.lock: - req_status.prefill_prompt_ids_event.prompt_ids = prompt_ids - req_status.prefill_prompt_ids_event.set() + req_status.prefill_prompt_ids_event.prompt_ids = prompt_ids + req_status.prefill_prompt_ids_event.set() except: logger.error( f"PD_UPLOAD_PREFILL_PROMPT_IDS fail find req status for group_req_id: {group_req_id}" @@ -652,26 +651,12 @@ def _split_max_new_tokens(self, max_new_tokens: int) -> List[int]: class ReqStatus: def __init__(self, req_id, p_node, d_node) -> None: self.req_id = req_id - self.lock = asyncio.Lock() - self.event = asyncio.Event() + self.out_tokens = AsyncQueue() self.up_status_event = asyncio.Event() self.prefill_prompt_ids_event = asyncio.Event() - self.out_token_info_list: List[Tuple[int, str, dict, FinishStatus]] = [] self.p_node: PD_Client_Obj = p_node self.d_node: PD_Client_Obj = d_node - async def drain_tokens(self, req_id_to_out_inf): - try: - await asyncio.wait_for(self.event.wait(), timeout=5) - except asyncio.TimeoutError: - pass - async with self.lock: - self.event.clear() - assert self.req_id in req_id_to_out_inf, f"error state req_id {self.req_id}" - tokens = self.out_token_info_list - self.out_token_info_list = [] - return tokens - class PDManager: def __init__(self, args: StartArgs): From aabe3b612255d2fae4bc6cbf439d0ef637c5bb6c Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 11 Aug 2026 12:06:45 +0800 Subject: [PATCH 12/21] style(pd): apply project formatting --- lightllm/server/httpserver_for_pd_master/manager.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 624971dc0..0d08d7af8 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -452,9 +452,7 @@ async def fetch_pd_stream( # P 首 token 携带缓存统计,需先于已缓冲的 D 输出返回。 if needs_prefill_first_token: - prefill_token = next( - (token for token in new_tokens if token[2].get("node_mode") == "prefill"), None - ) + prefill_token = next((token for token in new_tokens if token[2].get("node_mode") == "prefill"), None) if prefill_token is None: buffered_decode_tokens.extend(new_tokens) continue From 0f56003be02e6ccf36363a5250307233abee080c Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 11 Aug 2026 15:08:36 +0800 Subject: [PATCH 13/21] fix(pd): bound missing prefill token fallback --- .../httpserver_for_pd_master/manager.py | 35 ++++++---- .../test_pd_master_cached_tokens.py | 69 +++++++++++++++++++ 2 files changed, 90 insertions(+), 14 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 0d08d7af8..c7a3b250e 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -29,6 +29,8 @@ logger = init_logger(__name__) +_PREFILL_TOKEN_WAIT_TIMEOUT = 5 + class HttpServerManagerForPDMaster: def __init__( @@ -430,9 +432,13 @@ async def fetch_pd_stream( first_token_emitted = False needs_prefill_first_token = decode_node_info.ready_kv_len != len(prompt_ids) - 1 buffered_decode_tokens = [] + prefill_token_deadline = None prompt_cache_len_from_prefill = None while True: - new_tokens = await req_status.out_tokens.wait_to_get_all_data(timeout=5) + wait_timeout = 5 + if prefill_token_deadline is not None: + wait_timeout = min(wait_timeout, max(0, prefill_token_deadline - time.monotonic())) + new_tokens = await req_status.out_tokens.wait_to_get_all_data(timeout=wait_timeout) assert group_request_id in self.req_id_to_out_inf, f"error state req_id {group_request_id}" if await request.is_disconnected(): raise ClientDisconnected( @@ -440,26 +446,27 @@ async def fetch_pd_stream( reason="fetch_pd_stream decode period check network disconnected", ) - # 本轮无新 token 且 D 已完成时,释放缓冲避免因 P 首 token 缺失而挂起。 - if not new_tokens: - if not buffered_decode_tokens or not buffered_decode_tokens[-1][3].is_finished(): - continue - logger.warning(f"{group_request_id}: prefill token missing; releasing decode output") - for sub_req_id, request_output, metadata, finish_status in buffered_decode_tokens: - metadata.pop("node_mode", None) - yield sub_req_id, request_output, metadata, finish_status - return + if not new_tokens and not buffered_decode_tokens: + continue # P 首 token 携带缓存统计,需先于已缓冲的 D 输出返回。 if needs_prefill_first_token: prefill_token = next((token for token in new_tokens if token[2].get("node_mode") == "prefill"), None) if prefill_token is None: buffered_decode_tokens.extend(new_tokens) - continue - new_tokens.remove(prefill_token) - new_tokens = [prefill_token, *buffered_decode_tokens, *new_tokens] - buffered_decode_tokens.clear() + if prefill_token_deadline is None: + prefill_token_deadline = time.monotonic() + _PREFILL_TOKEN_WAIT_TIMEOUT + if time.monotonic() < prefill_token_deadline: + continue + logger.warning(f"{group_request_id}: prefill token missing; releasing decode output") + new_tokens = buffered_decode_tokens + buffered_decode_tokens = [] + else: + new_tokens.remove(prefill_token) + new_tokens = [prefill_token, *buffered_decode_tokens, *new_tokens] + buffered_decode_tokens.clear() needs_prefill_first_token = False + prefill_token_deadline = None for sub_req_id, request_output, metadata, finish_status in new_tokens: output_index = metadata.get("count_output_tokens") diff --git a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py index a3dc268a9..eccc5ab7d 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -1,10 +1,12 @@ import asyncio import copy +import pickle from types import SimpleNamespace import pytest from lightllm.server.core.objs import FinishStatus, SamplingParams +import lightllm.server.httpserver_for_pd_master.manager as pd_master_manager from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster @@ -85,3 +87,70 @@ def test_multi_block_keeps_first_block_hit(monkeypatch): cached = _collect(mgr, sp, monkeypatch, split=[3, 2]) assert cached[-1] == 30, cached assert mgr.recorded_cache_hit_rates == [pytest.approx(0.3)] + + +def test_missing_prefill_token_flushes_while_decode_is_still_running(monkeypatch): + group_request_id = 1 + clock = [0.0] + + def decode_token(index): + return ( + group_request_id, + str(index), + {"count_output_tokens": index, "node_mode": "decode", "prompt_cache_len": 0}, + FinishStatus(), + ) + + class OutputBatches: + def __init__(self): + self.batches = [(0.0, [decode_token(1)]), (6.0, [decode_token(2)]), (7.0, [decode_token(3)])] + + async def wait_to_get_all_data(self, timeout): + assert self.batches, "decode output was not released after the prefill-token deadline" + clock[0], batch = self.batches.pop(0) + return batch + + async def send_bytes(_): + pass + + manager = object.__new__(HttpServerManagerForPDMaster) + manager.args = SimpleNamespace(pd_node_id=7) + manager.req_id_to_out_inf = {} + p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=send_bytes)) + d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=send_bytes)) + + async def wait_for_stage(event, request, timeout, group_request_id, stage): + if stage == "prefill": + event.prompt_ids = [1, 2, 3, 4] + else: + decode_node_info = SimpleNamespace(ready_kv_len=0) + event.upkv_status = SimpleNamespace(pd_kv_trans_params=pickle.dumps(decode_node_info)) + manager.req_id_to_out_inf[group_request_id].out_tokens = OutputBatches() + + manager._wait_for_event_or_disconnect = wait_for_stage + monkeypatch.setattr(pd_master_manager.time, "monotonic", lambda: clock[0]) + + sampling_params = SamplingParams() + sampling_params.group_request_id = group_request_id + sampling_params.max_new_tokens = 2048 + + async def is_disconnected(): + return False + + async def collect(): + stream = manager.fetch_pd_stream( + p_node, + d_node, + prompt="prompt", + sampling_params=sampling_params, + multimodal_params=SimpleNamespace(), + request=SimpleNamespace(is_disconnected=is_disconnected), + ) + try: + return [await stream.__anext__() for _ in range(3)] + finally: + await stream.aclose() + + outputs = asyncio.run(collect()) + assert [output for _, output, _, _ in outputs] == ["1", "2", "3"] + assert all(not finish_status.is_finished() for _, _, _, finish_status in outputs) From 14448c35346f543aff884647fb6153f792bd5230 Mon Sep 17 00:00:00 2001 From: sufubao <47234901+sufubao@users.noreply.github.com> Date: Tue, 11 Aug 2026 15:15:17 +0800 Subject: [PATCH 14/21] Delete unit_tests/server/httpserver/test_pd_master_cached_tokens.py --- .../test_pd_master_cached_tokens.py | 69 ------------------- 1 file changed, 69 deletions(-) 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 eccc5ab7d..a3dc268a9 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -1,12 +1,10 @@ import asyncio import copy -import pickle from types import SimpleNamespace import pytest from lightllm.server.core.objs import FinishStatus, SamplingParams -import lightllm.server.httpserver_for_pd_master.manager as pd_master_manager from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster @@ -87,70 +85,3 @@ def test_multi_block_keeps_first_block_hit(monkeypatch): cached = _collect(mgr, sp, monkeypatch, split=[3, 2]) assert cached[-1] == 30, cached assert mgr.recorded_cache_hit_rates == [pytest.approx(0.3)] - - -def test_missing_prefill_token_flushes_while_decode_is_still_running(monkeypatch): - group_request_id = 1 - clock = [0.0] - - def decode_token(index): - return ( - group_request_id, - str(index), - {"count_output_tokens": index, "node_mode": "decode", "prompt_cache_len": 0}, - FinishStatus(), - ) - - class OutputBatches: - def __init__(self): - self.batches = [(0.0, [decode_token(1)]), (6.0, [decode_token(2)]), (7.0, [decode_token(3)])] - - async def wait_to_get_all_data(self, timeout): - assert self.batches, "decode output was not released after the prefill-token deadline" - clock[0], batch = self.batches.pop(0) - return batch - - async def send_bytes(_): - pass - - manager = object.__new__(HttpServerManagerForPDMaster) - manager.args = SimpleNamespace(pd_node_id=7) - manager.req_id_to_out_inf = {} - p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=send_bytes)) - d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=send_bytes)) - - async def wait_for_stage(event, request, timeout, group_request_id, stage): - if stage == "prefill": - event.prompt_ids = [1, 2, 3, 4] - else: - decode_node_info = SimpleNamespace(ready_kv_len=0) - event.upkv_status = SimpleNamespace(pd_kv_trans_params=pickle.dumps(decode_node_info)) - manager.req_id_to_out_inf[group_request_id].out_tokens = OutputBatches() - - manager._wait_for_event_or_disconnect = wait_for_stage - monkeypatch.setattr(pd_master_manager.time, "monotonic", lambda: clock[0]) - - sampling_params = SamplingParams() - sampling_params.group_request_id = group_request_id - sampling_params.max_new_tokens = 2048 - - async def is_disconnected(): - return False - - async def collect(): - stream = manager.fetch_pd_stream( - p_node, - d_node, - prompt="prompt", - sampling_params=sampling_params, - multimodal_params=SimpleNamespace(), - request=SimpleNamespace(is_disconnected=is_disconnected), - ) - try: - return [await stream.__anext__() for _ in range(3)] - finally: - await stream.aclose() - - outputs = asyncio.run(collect()) - assert [output for _, output, _, _ in outputs] == ["1", "2", "3"] - assert all(not finish_status.is_finished() for _, _, _, finish_status in outputs) From 59237368069e80307030c6616dcd3ccef5b78cf4 Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 13 Aug 2026 17:53:18 +0800 Subject: [PATCH 15/21] fix(pd): preserve terminal decode marker --- lightllm/server/httpserver_for_pd_master/manager.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index c7a3b250e..f4044135c 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -472,7 +472,8 @@ async def fetch_pd_stream( output_index = metadata.get("count_output_tokens") node_run_mode = metadata.pop("node_mode", None) if output_index == 1: - if first_token_emitted: + # D 首 token 可能是 KV 传输失败产生的唯一结束标记,不能按重复 token 丢弃。 + if first_token_emitted and not (node_run_mode == "decode" and finish_status.is_finished()): continue first_token_emitted = True if node_run_mode == "prefill": From 07de7a50ff587da8b9e34502b45253c5605f061c Mon Sep 17 00:00:00 2001 From: sufubao Date: Wed, 19 Aug 2026 17:08:03 +0800 Subject: [PATCH 16/21] fix(pd): serialize NIXL peer reconnects --- .../decode_node_impl/decode_trans_process.py | 1 - .../mode_backend/pd/nixl_kv_transporter.py | 415 +++++++++++------- .../prefill_trans_process.py | 6 +- .../test_nixl_kv_transporter_thread_safety.py | 249 +++++++++++ 4 files changed, 519 insertions(+), 152 deletions(-) create mode 100644 unit_tests/test_nixl_kv_transporter_thread_safety.py diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index cd9b015b1..af1efcd47 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -369,7 +369,6 @@ def request_page_loop(self): self.waiting_dict.pop(key, None) logger.error(f"send write ready task to prefill node failed: {trans_task.to_str()}") logger.exception(str(e)) - self.transporter.remove_remote_agent(peer_name=trans_task.prefill_agent_name) trans_task.error_info = f"send write ready task to prefill node failed: {str(e)}" self.failed_queue.put(trans_task) continue diff --git a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py index bd5e11f05..fa0512649 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py @@ -1,9 +1,9 @@ import pickle import copy import os +import threading import time -from dataclasses import dataclass -from typing import Dict +from typing import Dict, Optional, Tuple from torch import Tensor from lightllm.server.pd_io_struct import PDChunckedTransTask, PDAgentMetadata from lightllm.utils.log_utils import init_logger @@ -15,6 +15,7 @@ from nixl._api import nixl_agent as NixlWrapper from nixl._api import nixlBind from nixl._api import nixl_agent_config + from nixl._api import nixl_thread_sync_t logger.info("Nixl is available") except ImportError: @@ -26,20 +27,28 @@ class NixlKVTransporter: def __init__(self, node_id: int, tp_idx: int, kv_move_buffer: Tensor): self.node_id = node_id self.tp_idx = tp_idx + # A transporter is shared by the notification, transfer, status, and + # failure-handling threads. NIXL's internal synchronization protects + # individual calls; this lock also makes LightLLM's compound peer + # lifecycle operations atomic. + self._nixl_lock = threading.RLock() self.capture_telemetry = os.getenv("LIGHTLLM_NIXL_CAPTURE_TELEMETRY", "0").lower() in ( "1", "true", "yes", "on", ) - conf = None + conf = nixl_agent_config(sync_mode=nixl_thread_sync_t.NIXL_THREAD_SYNC_STRICT) if self.capture_telemetry: - conf = nixl_agent_config() conf.capture_telemetry = True logger.info("NIXL telemetry enabled") self.nixl_agent = NixlWrapper(self.agent_name, conf) self._register_kv_move_buffer(kv_move_buffer=kv_move_buffer) self.remote_agents: Dict[str, PDAgentMetadata] = {} + self._peer_generations: Dict[str, int] = {} + self._next_peer_generation = 0 + self._broken_remote_agents: set[str] = set() + self._active_xfers: Dict[int, Tuple[object, str, int]] = {} return @property @@ -48,14 +57,17 @@ def agent_name(self) -> str: @property def agent_metadata(self): - return self.nixl_agent.get_agent_metadata() + with self._nixl_lock: + return self.nixl_agent.get_agent_metadata() @property def local_page_mem_desc(self): - return self.nixl_agent.get_serialized_descs(self.page_reg_desc) + with self._nixl_lock: + return self.nixl_agent.get_serialized_descs(self.page_reg_desc) def get_new_notifs(self) -> Dict[str, list[bytes]]: - return self.nixl_agent.get_new_notifs() + with self._nixl_lock: + return self.nixl_agent.get_new_notifs() def _register_kv_move_buffer(self, kv_move_buffer: Tensor): self.num_pages, self.page_size, self.num_layers, self.kv_head_num, self.head_dims = kv_move_buffer.shape @@ -73,11 +85,54 @@ def _create_paged_xfer_handles(self, reg_desc: "nixlBind.nixlRegDList", page_num return self.nixl_agent.prep_xfer_dlist(agent_name, descs, "VRAM") def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): + with self._nixl_lock: + self._ensure_remote_agent_locked(remote_agent) + return + + @staticmethod + def _same_remote_agent(left: PDAgentMetadata, right: PDAgentMetadata) -> bool: + return ( + left.agent_name == right.agent_name + and left.agent_metadata == right.agent_metadata + and left.num_pages == right.num_pages + and left.page_reg_desc == right.page_reg_desc + ) + + def _ensure_remote_agent_locked(self, remote_agent: PDAgentMetadata) -> Tuple[PDAgentMetadata, int]: + peer_name = remote_agent.agent_name + current_agent = self.remote_agents.get(peer_name) + current_generation = self._peer_generations.get(peer_name) + + if current_agent is not None and not self._same_remote_agent(current_agent, remote_agent): + self._mark_remote_agent_broken_locked(peer_name, current_generation) + current_agent = self.remote_agents.get(peer_name) + + if peer_name in self._broken_remote_agents: + if current_agent is not None: + self._remove_remote_agent_locked(peer_name, expected_generation=current_generation) + else: + try: + self.nixl_agent.remove_remote_agent(peer_name) + self._broken_remote_agents.discard(peer_name) + except BaseException as e: + raise RuntimeError(f"NIXL remote agent {peer_name} is broken and could not be removed") from e + current_agent = self.remote_agents.get(peer_name) + + if current_agent is None: + self._connect_add_remote_agent_locked(remote_agent) + current_agent = self.remote_agents[peer_name] + current_generation = self._peer_generations[peer_name] + + if peer_name in self._broken_remote_agents: + raise RuntimeError(f"NIXL remote agent {peer_name} is unavailable while active transfers drain") + + return current_agent, current_generation + + def _connect_add_remote_agent_locked(self, remote_agent: PDAgentMetadata): if remote_agent.agent_name in self.remote_agents: return start_time = time.time() - peer_name = self.nixl_agent.add_remote_agent(remote_agent.agent_metadata) if isinstance(peer_name, bytes): peer_name = peer_name.decode() @@ -86,96 +141,135 @@ def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): peer_name == remote_agent.agent_name ), f"Peer name {peer_name} does not match remote name {remote_agent.agent_name}" - page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) - kv_page_xfer_handles = self._create_paged_xfer_handles( - page_mem_desc, remote_agent.num_pages, agent_name=peer_name - ) - remote_agent.page_xfer_handles = kv_page_xfer_handles + self._next_peer_generation += 1 + generation = self._next_peer_generation + self.remote_agents[peer_name] = remote_agent + self._peer_generations[peer_name] = generation + try: + page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) + remote_agent.page_xfer_handles = self._create_paged_xfer_handles( + page_mem_desc, remote_agent.num_pages, agent_name=peer_name + ) + except BaseException: + self._broken_remote_agents.add(peer_name) + self._remove_remote_agent_locked(peer_name, expected_generation=generation) + raise logger.info( - f"Added remote agent {peer_name} with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" + f"Added remote agent {peer_name} generation {generation} " + f"with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" ) - - self.remote_agents[remote_agent.agent_name] = remote_agent + self._broken_remote_agents.discard(peer_name) return def remove_remote_agent(self, peer_name: str): - if peer_name in self.remote_agents: + with self._nixl_lock: + generation = self._peer_generations.get(peer_name) + if generation is None: + logger.warning(f"try to remove remote agent, but peer name {peer_name} agent did not exist") + return + self._mark_remote_agent_broken_locked(peer_name, generation) + return + + def _has_active_xfers_locked(self, peer_name: str, generation: int) -> bool: + return any( + active_peer_name == peer_name and active_generation == generation + for _, active_peer_name, active_generation in self._active_xfers.values() + ) + + def _mark_remote_agent_broken_locked(self, peer_name: str, generation: Optional[int]): + if generation is None or self._peer_generations.get(peer_name) != generation: + return + self._broken_remote_agents.add(peer_name) + self._remove_remote_agent_locked(peer_name, expected_generation=generation) + + def _remove_remote_agent_locked(self, peer_name: str, expected_generation: Optional[int] = None) -> bool: + generation = self._peer_generations.get(peer_name) + if generation is None: + return False + if expected_generation is not None and generation != expected_generation: + return False + if self._has_active_xfers_locked(peer_name, generation): + self._broken_remote_agents.add(peer_name) + logger.warning( + f"defer removing remote agent {peer_name} generation {generation} until active transfers drain" + ) + return False + + remote_agent = self.remote_agents[peer_name] + try: + self.nixl_agent.remove_remote_agent(remote_agent.agent_name) + except BaseException as e: + self._broken_remote_agents.add(peer_name) + logger.error(f"remove remote agent {peer_name} generation {generation} failed") + logger.exception(str(e)) + return False + + self.remote_agents.pop(peer_name, None) + self._peer_generations.pop(peer_name, None) + self._broken_remote_agents.discard(peer_name) + if remote_agent.page_xfer_handles is not None: try: - remote_agent: PDAgentMetadata = self.remote_agents.pop(peer_name, None) - assert remote_agent.agent_name == peer_name - self.nixl_agent.remove_remote_agent(remote_agent.agent_name) - if remote_agent.page_xfer_handles is not None: - self.nixl_agent.release_dlist_handle(remote_agent.page_xfer_handles) + self.nixl_agent.release_dlist_handle(remote_agent.page_xfer_handles) except BaseException as e: - logger.error(f"remove remote agent {peer_name} failed") + logger.error(f"release remote agent {peer_name} descriptor handle failed") logger.exception(str(e)) - else: - logger.warning(f"try to remove remote agent, but peer name {peer_name} agent did not exist") + finally: + remote_agent.page_xfer_handles = None + return True + + def _send_notif_locked(self, peer_name: str, generation: int, notif_msg: bytes): + try: + self.nixl_agent.send_notif(remote_agent_name=peer_name, notif_msg=notif_msg) + except BaseException: + self._mark_remote_agent_broken_locked(peer_name, generation) + raise def send_write_done_task_to_decode_node(self, trans_task: PDChunckedTransTask): decode_agent_name = trans_task.decode_agent_name - if decode_agent_name not in self.remote_agents: - logger.warning(f"decode_agent_name {decode_agent_name} not exist") - _remote_agent = trans_task.create_decode_agent_obj() - self.connect_add_remote_agent(_remote_agent) - - new_trans_task: PDChunckedTransTask = copy.copy(trans_task) - new_trans_task.write_stage = "done" - new_trans_task.mem_indexes = None - new_trans_task.xfer_handle = None - new_trans_task.decode_agent_metadata = None - new_trans_task.decode_page_reg_desc = None - new_trans_task.prefill_agent_name = self.agent_name - new_trans_task.prefill_agent_metadata = self.agent_metadata - new_trans_task.prefill_num_pages = self.num_pages - new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc - self.nixl_agent.send_notif( - remote_agent_name=decode_agent_name, - notif_msg=pickle.dumps(new_trans_task), - ) + with self._nixl_lock: + _, generation = self._ensure_remote_agent_locked(trans_task.create_decode_agent_obj()) + new_trans_task: PDChunckedTransTask = copy.copy(trans_task) + new_trans_task.write_stage = "done" + new_trans_task.mem_indexes = None + new_trans_task.xfer_handle = None + new_trans_task.decode_agent_metadata = None + new_trans_task.decode_page_reg_desc = None + new_trans_task.prefill_agent_name = self.agent_name + new_trans_task.prefill_agent_metadata = self.agent_metadata + new_trans_task.prefill_num_pages = self.num_pages + new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc + self._send_notif_locked(decode_agent_name, generation, pickle.dumps(new_trans_task)) return def send_write_request_task_to_decode_node(self, trans_task: PDChunckedTransTask): decode_agent_name = trans_task.decode_agent_name - if decode_agent_name not in self.remote_agents: - logger.warning(f"decode_agent_name {decode_agent_name} not exist") - _remote_agent = trans_task.create_decode_agent_obj() - self.connect_add_remote_agent(_remote_agent) - - new_trans_task: PDChunckedTransTask = copy.copy(trans_task) - new_trans_task.write_stage = "request" - new_trans_task.mem_indexes = None - new_trans_task.xfer_handle = None - new_trans_task.prefill_agent_name = self.agent_name - new_trans_task.prefill_agent_metadata = self.agent_metadata - new_trans_task.prefill_num_pages = self.num_pages - new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc - self.nixl_agent.send_notif( - remote_agent_name=decode_agent_name, - notif_msg=pickle.dumps(new_trans_task), - ) + with self._nixl_lock: + _, generation = self._ensure_remote_agent_locked(trans_task.create_decode_agent_obj()) + new_trans_task: PDChunckedTransTask = copy.copy(trans_task) + new_trans_task.write_stage = "request" + new_trans_task.mem_indexes = None + new_trans_task.xfer_handle = None + new_trans_task.prefill_agent_name = self.agent_name + new_trans_task.prefill_agent_metadata = self.agent_metadata + new_trans_task.prefill_num_pages = self.num_pages + new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc + self._send_notif_locked(decode_agent_name, generation, pickle.dumps(new_trans_task)) return def send_write_ready_task_to_prefill_node(self, trans_task: PDChunckedTransTask): prefill_agent_name = trans_task.prefill_agent_name - if prefill_agent_name not in self.remote_agents: - logger.warning(f"prefill_agent_name {prefill_agent_name} not exist") - _remote_agent = trans_task.create_prefill_agent_obj() - self.connect_add_remote_agent(_remote_agent) - - new_trans_task: PDChunckedTransTask = copy.copy(trans_task) - new_trans_task.write_stage = "ready" - new_trans_task.mem_indexes = None - new_trans_task.xfer_handle = None - new_trans_task.decode_agent_name = self.agent_name - new_trans_task.decode_agent_metadata = self.agent_metadata - new_trans_task.decode_num_pages = self.num_pages - new_trans_task.decode_page_reg_desc = self.local_page_mem_desc - self.nixl_agent.send_notif( - remote_agent_name=prefill_agent_name, - notif_msg=pickle.dumps(new_trans_task), - ) + with self._nixl_lock: + _, generation = self._ensure_remote_agent_locked(trans_task.create_prefill_agent_obj()) + new_trans_task: PDChunckedTransTask = copy.copy(trans_task) + new_trans_task.write_stage = "ready" + new_trans_task.mem_indexes = None + new_trans_task.xfer_handle = None + new_trans_task.decode_agent_name = self.agent_name + new_trans_task.decode_agent_metadata = self.agent_metadata + new_trans_task.decode_num_pages = self.num_pages + new_trans_task.decode_page_reg_desc = self.local_page_mem_desc + self._send_notif_locked(prefill_agent_name, generation, pickle.dumps(new_trans_task)) return def send_error_info_to_prefill_node(self, trans_task: PDChunckedTransTask): @@ -186,53 +280,41 @@ def send_error_info_to_prefill_node(self, trans_task: PDChunckedTransTask): try: prefill_agent_name = trans_task.prefill_agent_name - if prefill_agent_name not in self.remote_agents: - logger.warning(f"prefill_agent_name {prefill_agent_name} not exist") - _remote_agent = trans_task.create_prefill_agent_obj() - self.connect_add_remote_agent(_remote_agent) - assert trans_task.error_info is not None - new_trans_task: PDChunckedTransTask = copy.copy(trans_task) - new_trans_task.write_stage = "error" - new_trans_task.mem_indexes = None - new_trans_task.xfer_handle = None - new_trans_task.decode_agent_name = self.agent_name - new_trans_task.decode_agent_metadata = self.agent_metadata - new_trans_task.decode_num_pages = self.num_pages - new_trans_task.decode_page_reg_desc = self.local_page_mem_desc - self.nixl_agent.send_notif( - remote_agent_name=prefill_agent_name, - notif_msg=pickle.dumps(new_trans_task), - ) + with self._nixl_lock: + _, generation = self._ensure_remote_agent_locked(trans_task.create_prefill_agent_obj()) + assert trans_task.error_info is not None + new_trans_task: PDChunckedTransTask = copy.copy(trans_task) + new_trans_task.write_stage = "error" + new_trans_task.mem_indexes = None + new_trans_task.xfer_handle = None + new_trans_task.decode_agent_name = self.agent_name + new_trans_task.decode_agent_metadata = self.agent_metadata + new_trans_task.decode_num_pages = self.num_pages + new_trans_task.decode_page_reg_desc = self.local_page_mem_desc + self._send_notif_locked(prefill_agent_name, generation, pickle.dumps(new_trans_task)) except BaseException as e: logger.error(f"send error info to prefill node failed: {trans_task.to_str()}") logger.exception(str(e)) - self.remove_remote_agent(peer_name=prefill_agent_name) return def send_error_info_to_decode_node(self, trans_task: PDChunckedTransTask): try: decode_agent_name = trans_task.decode_agent_name - if decode_agent_name not in self.remote_agents: - logger.warning(f"decode_agent_name {decode_agent_name} not exist") - _remote_agent = trans_task.create_decode_agent_obj() - self.connect_add_remote_agent(_remote_agent) - assert trans_task.error_info is not None - new_trans_task: PDChunckedTransTask = copy.copy(trans_task) - new_trans_task.write_stage = "error" - new_trans_task.mem_indexes = None - new_trans_task.xfer_handle = None - new_trans_task.prefill_agent_name = self.agent_name - new_trans_task.prefill_agent_metadata = self.agent_metadata - new_trans_task.prefill_num_pages = self.num_pages - new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc - self.nixl_agent.send_notif( - remote_agent_name=decode_agent_name, - notif_msg=pickle.dumps(new_trans_task), - ) + with self._nixl_lock: + _, generation = self._ensure_remote_agent_locked(trans_task.create_decode_agent_obj()) + assert trans_task.error_info is not None + new_trans_task: PDChunckedTransTask = copy.copy(trans_task) + new_trans_task.write_stage = "error" + new_trans_task.mem_indexes = None + new_trans_task.xfer_handle = None + new_trans_task.prefill_agent_name = self.agent_name + new_trans_task.prefill_agent_metadata = self.agent_metadata + new_trans_task.prefill_num_pages = self.num_pages + new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc + self._send_notif_locked(decode_agent_name, generation, pickle.dumps(new_trans_task)) except BaseException as e: logger.error(f"send error info to decode node failed: {trans_task.to_str()}") logger.exception(str(e)) - self.remove_remote_agent(peer_name=decode_agent_name) return def write_blocks_paged( @@ -243,46 +325,85 @@ def write_blocks_paged( prefill node call this function to write kv blocks into decode node pages """ decode_agent_name = trans_task.decode_agent_name - if decode_agent_name not in self.remote_agents: - logger.warning(f"decode_agent_name {decode_agent_name} not exist") - _remote_agent = trans_task.create_decode_agent_obj() - self.connect_add_remote_agent(_remote_agent) - - assert trans_task.src_page_index is not None and trans_task.dst_page_index is not None - remote_agent: PDAgentMetadata = self.remote_agents[decode_agent_name] - src_handle = self.page_local_xfer_handles - dst_handle = remote_agent.page_xfer_handles - handle = self.nixl_agent.make_prepped_xfer( - "WRITE", - src_handle, - [trans_task.src_page_index], - dst_handle, - [trans_task.dst_page_index], - b"", - ) - if not handle: - raise RuntimeError(f"make_prepped_xfer failed for task: {trans_task.to_str()}") - - self.nixl_agent.transfer(handle) - - return handle + with self._nixl_lock: + remote_agent, generation = self._ensure_remote_agent_locked(trans_task.create_decode_agent_obj()) + assert trans_task.src_page_index is not None and trans_task.dst_page_index is not None + handle = None + try: + handle = self.nixl_agent.make_prepped_xfer( + "WRITE", + self.page_local_xfer_handles, + [trans_task.src_page_index], + remote_agent.page_xfer_handles, + [trans_task.dst_page_index], + b"", + ) + if not handle: + raise RuntimeError(f"make_prepped_xfer failed for task: {trans_task.to_str()}") + self.nixl_agent.transfer(handle) + self._active_xfers[id(handle)] = (handle, decode_agent_name, generation) + return handle + except BaseException: + if handle: + try: + self.nixl_agent.release_xfer_handle(handle=handle) + except BaseException as release_error: + logger.error(f"release failed transfer handle for remote agent {decode_agent_name}") + logger.exception(str(release_error)) + self._mark_remote_agent_broken_locked(decode_agent_name, generation) + raise def check_task_status(self, trans_task: PDChunckedTransTask) -> str: assert trans_task.xfer_handle is not None handle = trans_task.xfer_handle - xfer_state = self.nixl_agent.check_xfer_state(handle) - if xfer_state == "ERR": - logger.warning(f"Transfer failed with trans task {trans_task.to_str()} for handle {handle}") - return xfer_state + with self._nixl_lock: + active_xfer = self._active_xfers.get(id(handle)) + try: + xfer_state = self.nixl_agent.check_xfer_state(handle) + except BaseException: + if active_xfer is not None: + _, peer_name, generation = active_xfer + self._mark_remote_agent_broken_locked(peer_name, generation) + raise + if xfer_state == "ERR": + logger.warning(f"Transfer failed with trans task {trans_task.to_str()} for handle {handle}") + if active_xfer is not None: + _, peer_name, generation = active_xfer + self._mark_remote_agent_broken_locked(peer_name, generation) + return xfer_state def release_xfer_handle(self, handle): - self.nixl_agent.release_xfer_handle(handle=handle) + with self._nixl_lock: + active_xfer = self._active_xfers.get(id(handle)) + self.nixl_agent.release_xfer_handle(handle=handle) + self._active_xfers.pop(id(handle), None) + if active_xfer is not None: + _, peer_name, generation = active_xfer + if peer_name in self._broken_remote_agents: + self._remove_remote_agent_locked(peer_name, expected_generation=generation) return + def get_xfer_telemetry(self, handle): + with self._nixl_lock: + return self.nixl_agent.get_xfer_telemetry(handle) + + def query_xfer_backend(self, handle): + with self._nixl_lock: + return self.nixl_agent.query_xfer_backend(handle) + def shutdown(self): - self.nixl_agent.deregister_memory(self.page_reg_desc) - self.nixl_agent.release_dlist_handle(self.page_local_xfer_handles) - agent_names = list(self.remote_agents.keys()) - for agent_name in agent_names: - self.remove_remote_agent(agent_name) + with self._nixl_lock: + for handle, _, _ in list(self._active_xfers.values()): + try: + self.nixl_agent.release_xfer_handle(handle=handle) + except BaseException as e: + logger.error("release active transfer handle during NIXL shutdown failed") + logger.exception(str(e)) + self._active_xfers.clear() + for agent_name in list(self.remote_agents.keys()): + generation = self._peer_generations.get(agent_name) + self._broken_remote_agents.add(agent_name) + self._remove_remote_agent_locked(agent_name, expected_generation=generation) + self.nixl_agent.deregister_memory(self.page_reg_desc) + self.nixl_agent.release_dlist_handle(self.page_local_xfer_handles) return diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py index 40bd2e42a..463f6a738 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py @@ -225,7 +225,6 @@ def ready_transfer_loop(self): logger.error(f"send WRITE request to decode failed: {trans_task.to_str()}") logger.exception(str(e)) trans_task.error_info = f"send WRITE request to decode failed: {str(e)}" - self.transporter.remove_remote_agent(peer_name=trans_task.decode_agent_name) self.failed_queue.put(trans_task) continue return @@ -316,7 +315,6 @@ def write_peer_kv_loop(self): except BaseException as e: logger.error(f"write_blocks_paged failed: {trans_task.to_str()}") logger.exception(str(e)) - self.transporter.remove_remote_agent(peer_name=trans_task.decode_agent_name) trans_task.error_info = f"write_blocks_paged failed: {str(e)}" self.failed_queue.put(trans_task) continue @@ -342,11 +340,11 @@ def update_task_status_loop( if ret == "DONE": trans_task = self.waiting_dict.pop(trans_task.get_key(), None) if self.transporter.capture_telemetry: - telem = self.transporter.nixl_agent.get_xfer_telemetry(trans_task.xfer_handle) + telem = self.transporter.get_xfer_telemetry(trans_task.xfer_handle) total_us = telem.xferDuration post_us = telem.postDuration backend_us = telem.xferDuration - telem.postDuration - nixl_backend = self.transporter.nixl_agent.query_xfer_backend(trans_task.xfer_handle) + nixl_backend = self.transporter.query_xfer_backend(trans_task.xfer_handle) logger.info( f"write trans task request_id={trans_task.request_id} " f"kv=[{trans_task.start_kv_index},{trans_task.end_kv_index}) " diff --git a/unit_tests/test_nixl_kv_transporter_thread_safety.py b/unit_tests/test_nixl_kv_transporter_thread_safety.py new file mode 100644 index 000000000..f32ab8852 --- /dev/null +++ b/unit_tests/test_nixl_kv_transporter_thread_safety.py @@ -0,0 +1,249 @@ +import threading +import time + +import pytest + +from lightllm.server.pd_io_struct import PDChunckedTransTask, PDAgentMetadata +from lightllm.server.router.model_infer.mode_backend.pd import nixl_kv_transporter + + +class _FakeBuffer: + shape = (1, 1, 1, 1, 1) + + @staticmethod + def element_size(): + return 1 + + +class _FakeConfig: + def __init__(self, sync_mode=None): + self.sync_mode = sync_mode + self.capture_telemetry = False + + +class _FakeSyncMode: + NIXL_THREAD_SYNC_STRICT = object() + + +class _FakeNixlAgent: + instances = [] + + def __init__(self, name, config): + self.name = name + self.config = config + self.add_calls = [] + self.remove_calls = [] + self.released_dlists = [] + self.released_xfers = [] + self.send_calls = [] + self.fail_next_send = False + self._native_call_lock = threading.Lock() + self._native_calls = 0 + self.max_native_calls = 0 + self.__class__.instances.append(self) + + def _enter_native_call(self): + with self._native_call_lock: + self._native_calls += 1 + self.max_native_calls = max(self.max_native_calls, self._native_calls) + + def _leave_native_call(self): + with self._native_call_lock: + self._native_calls -= 1 + + def register_memory(self, _buffer): + return [(1000, 1, 0, "")] + + def get_xfer_descs(self, pages_data, _mem_type): + return pages_data + + def prep_xfer_dlist(self, agent_name, descs, _mem_type): + return (agent_name, tuple(descs)) + + def get_agent_metadata(self): + return b"local-metadata" + + def get_serialized_descs(self, _reg_desc): + return b"local-page-desc" + + def get_new_notifs(self): + return {} + + def add_remote_agent(self, metadata): + self._enter_native_call() + try: + time.sleep(0.01) + self.add_calls.append(metadata) + return metadata.decode() + finally: + self._leave_native_call() + + def deserialize_descs(self, _page_reg_desc): + return [(2000, 1, 0, "")] + + def remove_remote_agent(self, peer_name): + self._enter_native_call() + try: + self.remove_calls.append(peer_name) + finally: + self._leave_native_call() + + def release_dlist_handle(self, handle): + self.released_dlists.append(handle) + + def send_notif(self, remote_agent_name, notif_msg): + self._enter_native_call() + try: + self.send_calls.append((remote_agent_name, notif_msg)) + if self.fail_next_send: + self.fail_next_send = False + raise RuntimeError("remote disconnected") + finally: + self._leave_native_call() + + def make_prepped_xfer(self, *_args): + return object() + + def transfer(self, _handle): + return "PROC" + + def check_xfer_state(self, _handle): + return "DONE" + + def release_xfer_handle(self, handle): + self.released_xfers.append(handle) + + def deregister_memory(self, _reg_desc): + return None + + +@pytest.fixture +def transporter(monkeypatch): + _FakeNixlAgent.instances.clear() + monkeypatch.setattr(nixl_kv_transporter, "NixlWrapper", _FakeNixlAgent) + monkeypatch.setattr(nixl_kv_transporter, "nixl_agent_config", _FakeConfig, raising=False) + monkeypatch.setattr(nixl_kv_transporter, "nixl_thread_sync_t", _FakeSyncMode, raising=False) + return nixl_kv_transporter.NixlKVTransporter(node_id=1, tp_idx=0, kv_move_buffer=_FakeBuffer()) + + +def _remote_agent(name="peer"): + return PDAgentMetadata( + agent_name=name, + agent_metadata=name.encode(), + num_pages=1, + page_reg_desc=b"remote-page-desc", + ) + + +def _task(remote_agent): + return PDChunckedTransTask( + request_id=1, + start_kv_index=0, + end_kv_index=1, + time_out_secs=60, + pd_master_node_id=1, + prefill_dp_index=0, + decode_dp_index=0, + src_device_id=0, + dst_device_id=0, + mem_indexes=[0], + prefill_agent_name=remote_agent.agent_name, + prefill_agent_metadata=remote_agent.agent_metadata, + prefill_num_pages=remote_agent.num_pages, + prefill_page_reg_desc=remote_agent.page_reg_desc, + decode_agent_name="decode", + decode_agent_metadata=b"decode", + decode_num_pages=1, + decode_page_reg_desc=b"decode-page-desc", + first_gen_token_id=None, + first_gen_token_logprob=None, + src_page_index=0, + dst_page_index=0, + ) + + +def _run_concurrently(callables): + barrier = threading.Barrier(len(callables) + 1) + errors = [] + + def run(callable_): + barrier.wait() + try: + callable_() + except BaseException as exc: + errors.append(exc) + + threads = [threading.Thread(target=run, args=(callable_,)) for callable_ in callables] + for thread in threads: + thread.start() + barrier.wait() + for thread in threads: + thread.join(timeout=2) + assert not thread.is_alive() + return errors + + +def test_uses_strict_nixl_sync_mode(transporter): + assert transporter.nixl_agent.config.sync_mode is _FakeSyncMode.NIXL_THREAD_SYNC_STRICT + + +def test_concurrent_peer_admission_calls_native_add_once(transporter): + remote_agent = _remote_agent() + + errors = _run_concurrently( + [ + lambda: transporter.connect_add_remote_agent(remote_agent), + lambda: transporter.connect_add_remote_agent(remote_agent), + ] + ) + + assert errors == [] + assert transporter.nixl_agent.add_calls == [b"peer"] + assert transporter.nixl_agent.max_native_calls == 1 + + +def test_failed_send_is_removed_before_one_serialized_reconnect(transporter): + remote_agent = _remote_agent() + task = _task(remote_agent) + transporter.connect_add_remote_agent(remote_agent) + transporter.nixl_agent.fail_next_send = True + + errors = _run_concurrently( + [ + lambda: transporter.send_write_ready_task_to_prefill_node(task), + lambda: transporter.send_write_ready_task_to_prefill_node(task), + ] + ) + + assert len(errors) == 1 + assert str(errors[0]) == "remote disconnected" + assert transporter.nixl_agent.add_calls == [b"peer", b"peer"] + assert transporter.nixl_agent.remove_calls == ["peer"] + assert transporter._peer_generations == {"peer": 2} + assert transporter.nixl_agent.max_native_calls == 1 + + +def test_active_transfer_defers_removal_and_stale_generation_cannot_remove_reconnect(transporter): + remote_agent = _remote_agent("decode") + task = _task(_remote_agent()) + handle = transporter.write_blocks_paged(task) + first_generation = transporter._peer_generations["decode"] + + with transporter._nixl_lock: + transporter._mark_remote_agent_broken_locked("decode", first_generation) + + assert "decode" in transporter.remote_agents + assert "decode" in transporter._broken_remote_agents + with pytest.raises(RuntimeError, match="active transfers drain"): + transporter.connect_add_remote_agent(remote_agent) + + transporter.release_xfer_handle(handle) + transporter.connect_add_remote_agent(remote_agent) + second_generation = transporter._peer_generations["decode"] + assert second_generation > first_generation + + with transporter._nixl_lock: + transporter._mark_remote_agent_broken_locked("decode", first_generation) + + assert transporter._peer_generations == {"decode": second_generation} + assert "decode" in transporter.remote_agents From 348aa8802fcc448d6b433d59ba66eeefb71fe16b Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 20 Aug 2026 22:26:47 +0800 Subject: [PATCH 17/21] fix(pd): cancel requests waiting for decode slots --- lightllm/server/httpserver/manager.py | 63 ++++++++++------ lightllm/server/httpserver/pd_loop.py | 49 ++++++++++--- .../httpserver/test_pd_pending_abort.py | 72 +++++++++++++++++++ 3 files changed, 151 insertions(+), 33 deletions(-) create mode 100644 unit_tests/server/httpserver/test_pd_pending_abort.py diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 195ef17e4..edd2b192b 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -311,6 +311,44 @@ def alloc_req_id(self, sampling_params): assert False, "dead code path" return group_request_id + async def _alloc_req_objs( + self, group_request_id, prompt_ids, sampling_params + ) -> List[Req]: + """Allocate and initialize request slots without leaking partial allocations.""" + alloced_req_indexes = [] + req_objs = [] + try: + 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) + + for i, req_index in enumerate(alloced_req_indexes): + req_obj = await self.shm_req_manager.async_get_req_obj_by_index( + req_index + ) + req_objs.append(req_obj) + req_obj.init( + group_request_id + i, + prompt_ids, + sampling_params, + self.tokenizer, + chunked_prefill_size=self.args.chunked_prefill_size, + ) + return req_objs + except BaseException: + for req_obj in req_objs: + await self.shm_req_manager.async_put_back_req_obj(req_obj) + for req_index in alloced_req_indexes: + await self.shm_req_manager.async_release_req_index(req_index) + raise + async def generate( self, prompt: Union[str, List[int]], @@ -414,28 +452,9 @@ async def generate( raise PDPrefillNodeStopGenToken(group_request_id=group_request_id) # 申请资源并存储 - 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) - 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) - req_obj.init( - group_request_id + i, - prompt_ids, - sampling_params, - self.tokenizer, - chunked_prefill_size=self.args.chunked_prefill_size, - ) - req_objs.append(req_obj) + req_objs = await self._alloc_req_objs( + group_request_id, prompt_ids, sampling_params + ) self._log_stage_timing( group_request_id, start_time, diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index a69364c95..8ab2f7f0f 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -34,6 +34,25 @@ async def timer_log(manager: HttpServerManager): return +async def _abort_pd_request( + manager: HttpServerManager, + group_req_id: int, + generation_tasks: Dict[int, asyncio.Task], +) -> bool: + generation_task = generation_tasks.get(group_req_id) + if ( + generation_task is not None + and not generation_task.done() + and not generation_task.cancelling() + ): + # The request may still be waiting for a shared-memory slot and therefore + # not yet be visible to HttpServerManager.abort(). Cancelling its owning + # task makes the abort effective at every point in the admission path. + generation_task.cancel() + + return (await manager.abort(group_req_id)) or generation_task is not None + + 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") @@ -84,6 +103,7 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O while True: forwarding_tokens_task = None heartbeat_task = None + generation_tasks: Dict[int, asyncio.Task] = {} try: uri = f"ws://{pd_master_obj.host_ip_port}/pd_register" async with websockets.connect( @@ -124,7 +144,7 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O group_req_id = sampling_params.group_request_id pd_event = asyncio.Event() group_req_id_to_event[group_req_id] = pd_event - asyncio.create_task( + generation_task = asyncio.create_task( _pd_process_generate( manager=manager, prompt=prompt, @@ -135,18 +155,18 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O pd_event=pd_event, ) ) + generation_tasks[group_req_id] = generation_task + + def remove_generation_task(done_task, request_id=group_req_id): + if generation_tasks.get(request_id) is done_task: + generation_tasks.pop(request_id, None) + + generation_task.add_done_callback(remove_generation_task) elif obj[0] == ObjType.ABORT: group_req_id = obj[1] logger.warning(f"recv cmd aborted req id {group_req_id}") - if not (await manager.abort(group_req_id)): - - async def delayed_abort_task(group_req_id, retry_count): - for _ in range(retry_count): - await asyncio.sleep(5.0) - if await manager.abort(group_req_id): - break - - asyncio.create_task(delayed_abort_task(group_req_id=group_req_id, retry_count=4)) + group_req_id_to_event.pop(group_req_id, None) + await _abort_pd_request(manager, group_req_id, generation_tasks) elif obj[0] == ObjType.PD_REQ_DECODE_NODE_INFO: _, group_req_id, decode_node_info = obj @@ -169,8 +189,10 @@ async def delayed_abort_task(group_req_id, retry_count): logger.exception(str(e)) finally: child_tasks = [task for task in (forwarding_tokens_task, heartbeat_task) if task is not None] + child_tasks.extend(generation_tasks.values()) for task in child_tasks: - task.cancel() + if not task.done() and not task.cancelling(): + task.cancel() if child_tasks: await asyncio.gather(*child_tasks, return_exceptions=True) @@ -235,6 +257,11 @@ async def _pd_process_generate( 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 asyncio.CancelledError: + logger.info( + f"pd request task cancelled for group_request_id {sampling_params.group_request_id}" + ) + raise except BaseException as e: logger.error(str(e)) diff --git a/unit_tests/server/httpserver/test_pd_pending_abort.py b/unit_tests/server/httpserver/test_pd_pending_abort.py new file mode 100644 index 000000000..098744ac3 --- /dev/null +++ b/unit_tests/server/httpserver/test_pd_pending_abort.py @@ -0,0 +1,72 @@ +import asyncio +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from lightllm.server.httpserver.manager import HttpServerManager +from lightllm.server.httpserver.pd_loop import _abort_pd_request + + +def test_pd_abort_cancels_request_before_http_manager_registration(): + async def run(): + manager = MagicMock() + manager.abort = AsyncMock(return_value=False) + request_started = asyncio.Event() + + async def wait_for_request_slot(): + request_started.set() + await asyncio.Future() + + generation_task = asyncio.create_task(wait_for_request_slot()) + generation_tasks = {123: generation_task} + await request_started.wait() + + assert await _abort_pd_request(manager, 123, generation_tasks) + assert generation_task.cancelling() == 1 + assert await _abort_pd_request(manager, 123, generation_tasks) + assert generation_task.cancelling() == 1 + with pytest.raises(asyncio.CancelledError): + await generation_task + + assert manager.abort.await_count == 2 + manager.abort.assert_awaited_with(123) + assert generation_task.cancelled() + + asyncio.run(run()) + + +def test_cancelled_slot_allocation_releases_partially_allocated_indexes(): + async def run(): + manager = HttpServerManager.__new__(HttpServerManager) + manager.shm_req_manager = MagicMock() + manager.shm_req_manager.async_release_req_index = AsyncMock() + manager.shm_req_manager.async_put_back_req_obj = AsyncMock() + manager.tokenizer = MagicMock() + manager.args = MagicMock(chunked_prefill_size=16) + + waiting_for_second_slot = asyncio.Event() + allocation_count = 0 + + async def alloc_req_index(): + nonlocal allocation_count + allocation_count += 1 + if allocation_count == 1: + return 7 + waiting_for_second_slot.set() + return None + + manager.shm_req_manager.async_alloc_req_index = alloc_req_index + sampling_params = MagicMock(n=2) + allocation_task = asyncio.create_task( + manager._alloc_req_objs(123, [1, 2], sampling_params) + ) + await waiting_for_second_slot.wait() + allocation_task.cancel() + + with pytest.raises(asyncio.CancelledError): + await allocation_task + + manager.shm_req_manager.async_release_req_index.assert_awaited_once_with(7) + manager.shm_req_manager.async_put_back_req_obj.assert_not_awaited() + + asyncio.run(run()) From f6a441755f2589c09bdbff8747a645c3ad397d6c Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 20 Aug 2026 23:11:20 +0800 Subject: [PATCH 18/21] fix(pd): fail requests when assigned nodes disconnect Track registrations by their owning WebSocket so a stale disconnect cannot evict a replacement connection. Wake requests assigned to a replaced or disconnected node and surface the connection failure at every PD wait stage. --- lightllm/server/api_http_pd.py | 2 +- .../httpserver_for_pd_master/manager.py | 52 ++++++++++++-- unit_tests/server/test_pd_master_mode.py | 70 ++++++++++++++++++- 3 files changed, 116 insertions(+), 8 deletions(-) diff --git a/lightllm/server/api_http_pd.py b/lightllm/server/api_http_pd.py index 1d8b2112f..c23e5e7a7 100644 --- a/lightllm/server/api_http_pd.py +++ b/lightllm/server/api_http_pd.py @@ -57,7 +57,7 @@ async def register_and_keep_alive(websocket: WebSocket): logger.exception(str(e)) finally: logger.error(f"client {regist_json} removed") - await g_objs.httpserver_manager.remove_pd(regist_json) + await g_objs.httpserver_manager.remove_pd(regist_json, websocket) return diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index e5bddd746..550706bf5 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -76,11 +76,26 @@ def is_healthy(self): return False async def register_pd(self, pd_info_json, websocket): - self.pd_manager.register_pd(pd_info_json, websocket) + replaced_node = self.pd_manager.register_pd(pd_info_json, websocket) + if replaced_node is not None: + await self._fail_requests_for_node(replaced_node, "PD node connection was replaced") return - async def remove_pd(self, pd_info_json): - self.pd_manager.remove_pd(pd_info_json) + async def remove_pd(self, pd_info_json, websocket): + removed_node = self.pd_manager.remove_pd(pd_info_json, websocket) + if removed_node is not None: + await self._fail_requests_for_node(removed_node, "PD node disconnected") + return + + async def _fail_requests_for_node(self, node: PD_Client_Obj, reason: str): + error_message = f"{reason}: {node.mode} node {node.client_ip_port}" + affected_requests = [ + req_status + for req_status in list(self.req_id_to_out_inf.values()) + if req_status.p_node is node or req_status.d_node is node + ] + for req_status in affected_requests: + await req_status.fail(RuntimeError(error_message)) return async def update_req_status(self, upkv_status: PDUpKVStatus): @@ -397,6 +412,7 @@ async def fetch_pd_stream( group_request_id=group_request_id, stage="prefill", ) + req_status.raise_if_failed() except ServerBusyError: logger.warning(f"group_request_id: {group_request_id} wait prefill prompt ids time out") raise @@ -417,6 +433,7 @@ async def fetch_pd_stream( group_request_id=group_request_id, stage="decode", ) + req_status.raise_if_failed() except ServerBusyError: logger.warning(f"group_request_id: {group_request_id} wait decode stage time out err, server is busy now.") raise @@ -439,6 +456,7 @@ async def fetch_pd_stream( if prefill_token_deadline is not None: wait_timeout = min(wait_timeout, max(0, prefill_token_deadline - time.monotonic())) new_tokens = await req_status.out_tokens.wait_to_get_all_data(timeout=wait_timeout) + req_status.raise_if_failed() assert group_request_id in self.req_id_to_out_inf, f"error state req_id {group_request_id}" if await request.is_disconnected(): raise ClientDisconnected( @@ -662,6 +680,21 @@ def __init__(self, req_id, p_node, d_node) -> None: self.prefill_prompt_ids_event = asyncio.Event() self.p_node: PD_Client_Obj = p_node self.d_node: PD_Client_Obj = d_node + self.error: Optional[BaseException] = None + + async def fail(self, error: BaseException): + async with self.out_tokens.lock: + if self.error is None: + self.error = error + self.out_tokens.event.set() + self.up_status_event.set() + self.prefill_prompt_ids_event.set() + return + + def raise_if_failed(self): + if self.error is not None: + raise self.error + return class PDManager: @@ -725,6 +758,7 @@ async def check_pd_nodes_health(self): def register_pd(self, pd_info_json, websocket): pd_client = PD_Client_Obj(**pd_info_json) + replaced_node = self.url_to_pd_nodes.get(pd_client.client_ip_port) 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( @@ -761,10 +795,16 @@ def register_pd(self, pd_info_json, websocket): self.selector.update_nodes(self.prefill_nodes, self.decode_nodes) logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} registed") - return + if replaced_node is not None and replaced_node.websocket is not websocket: + return replaced_node + return None - def remove_pd(self, pd_info_json): + def remove_pd(self, pd_info_json, websocket): pd_client = PD_Client_Obj(**pd_info_json) + registered_node = self.url_to_pd_nodes.get(pd_client.client_ip_port) + if registered_node is None or registered_node.websocket is not websocket: + logger.info(f"ignore stale disconnect for {pd_client.mode} node {pd_client.client_ip_port}") + return None 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] @@ -773,7 +813,7 @@ def remove_pd(self, pd_info_json): self.selector.update_nodes(self.prefill_nodes, self.decode_nodes) logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} removed") - return + return registered_node def update_node_load_info(self, load_info: Optional[dict]): """更新节点负载信息 diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index edb758a94..afef60108 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -6,7 +6,20 @@ from easydict import EasyDict from lightllm.server.core.objs.start_args_type import StartArgs -from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager +from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager, ReqStatus + + +def _prefill_pd_info(args, node_id=1, client_ip_port="10.0.0.1:8000"): + return { + "node_id": node_id, + "client_ip_port": client_ip_port, + "mode": "prefill", + "start_args": { + "max_req_total_len": args.max_req_total_len, + "max_image_pixels": args.max_image_pixels, + "disable_image_resize": args.disable_image_resize, + }, + } def test_auto_set_response_parsers_from_qwen35_model_config(tmp_path): @@ -243,6 +256,61 @@ def pd_info(node_id, client_ip_port): assert [node.dispatched_req_num for node in manager.prefill_nodes] == [1, 0] +def test_pd_reconnection_fails_old_requests_and_ignores_stale_disconnect(): + async def run(): + args = StartArgs() + pd_manager = PDManager(args) + http_manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + http_manager.pd_manager = pd_manager + http_manager.req_id_to_out_inf = {} + pd_info = _prefill_pd_info(args) + old_websocket = object() + new_websocket = object() + + pd_manager.register_pd(pd_info, old_websocket) + old_node = pd_manager.prefill_nodes[0] + req_status = ReqStatus(8, old_node, object()) + http_manager.req_id_to_out_inf[8] = req_status + + await http_manager.register_pd(pd_info, new_websocket) + + with pytest.raises(RuntimeError, match="connection was replaced"): + req_status.raise_if_failed() + assert req_status.prefill_prompt_ids_event.is_set() + assert req_status.up_status_event.is_set() + assert req_status.out_tokens.event.is_set() + + await http_manager.remove_pd(pd_info, old_websocket) + assert pd_manager.url_to_pd_nodes[pd_info["client_ip_port"]].websocket is new_websocket + + asyncio.run(run()) + + +def test_pd_disconnect_wakes_requests_owned_by_the_disconnected_node(): + async def run(): + args = StartArgs() + pd_manager = PDManager(args) + http_manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + http_manager.pd_manager = pd_manager + http_manager.req_id_to_out_inf = {} + pd_info = _prefill_pd_info(args) + websocket = object() + + pd_manager.register_pd(pd_info, websocket) + node = pd_manager.prefill_nodes[0] + req_status = ReqStatus(9, node, object()) + http_manager.req_id_to_out_inf[9] = req_status + token_waiter = asyncio.create_task(req_status.out_tokens.wait_to_get_all_data(timeout=10)) + + await http_manager.remove_pd(pd_info, websocket) + assert await asyncio.wait_for(token_waiter, timeout=1) == [] + with pytest.raises(RuntimeError, match="PD node disconnected"): + req_status.raise_if_failed() + assert pd_info["client_ip_port"] not in pd_manager.url_to_pd_nodes + + asyncio.run(run()) + + def test_pd_master_inference_health_matches_normal_node_semantics(monkeypatch): monkeypatch.setattr("lightllm.server.httpserver_for_pd_master.manager.time.time", lambda: 1000) manager = SimpleNamespace( From 623d640daf4305a39b339f149f80398d03892507 Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 20 Aug 2026 23:11:24 +0800 Subject: [PATCH 19/21] fix(pd): propagate worker background failures Supervise token forwarding, heartbeat, and per-request generation tasks while the worker waits for PD master messages. Tear down the connection when a background task fails so the master observes a disconnect instead of waiting on a silently dead request. --- lightllm/server/httpserver/pd_loop.py | 54 ++++++++++++++-- .../httpserver/test_pd_connection_tasks.py | 63 +++++++++++++++++++ 2 files changed, 112 insertions(+), 5 deletions(-) create mode 100644 unit_tests/server/httpserver/test_pd_connection_tasks.py diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index 8ab2f7f0f..697bb65e1 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -53,6 +53,33 @@ async def _abort_pd_request( return (await manager.abort(group_req_id)) or generation_task is not None +async def _recv_or_raise_on_background_failure(websocket, background_tasks): + recv_task = asyncio.create_task(websocket.recv()) + try: + done, _ = await asyncio.wait( + [recv_task, *background_tasks], + return_when=asyncio.FIRST_COMPLETED, + ) + for task in background_tasks: + if task not in done: + continue + + recv_task.cancel() + await asyncio.gather(recv_task, return_exceptions=True) + if task.cancelled(): + raise asyncio.CancelledError + error = task.exception() + if error is not None: + raise error + raise RuntimeError("PD connection background task exited unexpectedly") + + return recv_task.result() + finally: + if not recv_task.done(): + recv_task.cancel() + await asyncio.gather(recv_task, return_exceptions=True) + + 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") @@ -103,6 +130,7 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O while True: forwarding_tokens_task = None heartbeat_task = None + connection_failure = None generation_tasks: Dict[int, asyncio.Task] = {} try: uri = f"ws://{pd_master_obj.host_ip_port}/pd_register" @@ -133,11 +161,15 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O # 转发任务 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)) + connection_failure = asyncio.get_running_loop().create_future() group_req_id_to_event: Dict[int, asyncio.Event] = weakref.WeakValueDictionary() # 接收 pd master 发来的请求,并推理后,将生成的token转发回pd master。 while True: - recv_bytes = await websocket.recv() + recv_bytes = await _recv_or_raise_on_background_failure( + websocket, + (forwarding_tokens_task, heartbeat_task, connection_failure), + ) obj = pickle.loads(recv_bytes) if obj[0] == ObjType.REQ: prompt, sampling_params, multimodal_params = obj[1] @@ -160,6 +192,10 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O def remove_generation_task(done_task, request_id=group_req_id): if generation_tasks.get(request_id) is done_task: generation_tasks.pop(request_id, None) + if not done_task.cancelled(): + error = done_task.exception() + if error is not None and not connection_failure.done(): + connection_failure.set_exception(error) generation_task.add_done_callback(remove_generation_task) elif obj[0] == ObjType.ABORT: @@ -190,9 +226,14 @@ def remove_generation_task(done_task, request_id=group_req_id): finally: child_tasks = [task for task in (forwarding_tokens_task, heartbeat_task) if task is not None] child_tasks.extend(generation_tasks.values()) + if connection_failure is not None: + child_tasks.append(connection_failure) for task in child_tasks: - if not task.done() and not task.cancelling(): - task.cancel() + if task.done(): + continue + if isinstance(task, asyncio.Task) and task.cancelling(): + continue + task.cancel() if child_tasks: await asyncio.gather(*child_tasks, return_exceptions=True) @@ -262,8 +303,11 @@ async def _pd_process_generate( f"pd request task cancelled for group_request_id {sampling_params.group_request_id}" ) raise - except BaseException as e: - logger.error(str(e)) + except BaseException: + logger.exception( + f"PD request generation failed for group_request_id {sampling_params.group_request_id}" + ) + raise # 转发token的task diff --git a/unit_tests/server/httpserver/test_pd_connection_tasks.py b/unit_tests/server/httpserver/test_pd_connection_tasks.py new file mode 100644 index 000000000..6a4fd8367 --- /dev/null +++ b/unit_tests/server/httpserver/test_pd_connection_tasks.py @@ -0,0 +1,63 @@ +import asyncio +from types import SimpleNamespace + +import pytest + +from lightllm.server.httpserver.async_queue import AsyncQueue +from lightllm.server.httpserver.pd_loop import ( + _pd_process_generate, + _recv_or_raise_on_background_failure, +) + + +def test_background_failure_interrupts_blocked_pd_receive(): + async def run(): + recv_started = asyncio.Event() + recv_cancelled = asyncio.Event() + + class BlockingWebsocket: + async def recv(self): + recv_started.set() + try: + await asyncio.Future() + except asyncio.CancelledError: + recv_cancelled.set() + raise + + websocket = BlockingWebsocket() + failure = asyncio.get_running_loop().create_future() + receive_task = asyncio.create_task( + _recv_or_raise_on_background_failure(websocket, (failure,)) + ) + await recv_started.wait() + + failure.set_exception(RuntimeError("generation failed")) + with pytest.raises(RuntimeError, match="generation failed"): + await receive_task + assert recv_cancelled.is_set() + + asyncio.run(run()) + + +def test_pd_generation_failure_is_not_swallowed(): + async def run(): + class FailingManager: + args = SimpleNamespace(run_mode="prefill") + + async def generate(self, **_kwargs): + yield 1, "token", {}, None + raise RuntimeError("generation failed") + + sampling_params = SimpleNamespace(group_request_id=123) + with pytest.raises(RuntimeError, match="generation failed"): + await _pd_process_generate( + manager=FailingManager(), + prompt="prompt", + sampling_params=sampling_params, + multimodal_params={}, + forwarding_queue=AsyncQueue(), + pd_upload_websocket=object(), + pd_event=asyncio.Event(), + ) + + asyncio.run(run()) From feedc43e5455e291263b61359a967dfbe0e9dced Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 20 Aug 2026 23:11:30 +0800 Subject: [PATCH 20/21] fix(pd): recover NIXL transfer task failures Route status-query and completion-notification errors through the existing failed-task cleanup path so pages and transfer handles are reclaimed. Install a fatal thread exception hook in transfer subprocesses so any unhandled daemon-thread failure is visible to process supervision. --- .../decode_node_impl/decode_trans_process.py | 3 +- .../prefill_trans_process.py | 93 ++++++++++++------- lightllm/utils/process_check.py | 15 +++ .../test_nixl_kv_transporter_thread_safety.py | 72 ++++++++++++++ 4 files changed, 150 insertions(+), 33 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index af1efcd47..7373671ab 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -21,7 +21,7 @@ from ..kv_transporter import create_kv_transporter from lightllm.utils.error_utils import log_exception from lightllm.utils.envs_utils import get_unique_server_name -from lightllm.utils.process_check import start_parent_check_thread +from lightllm.utils.process_check import install_fatal_thread_excepthook, start_parent_check_thread logger = init_logger(__name__) @@ -47,6 +47,7 @@ def _init_env( task_out_queue: mp.Queue, up_status_in_queue: Optional[mp.SimpleQueue], ): + install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py index 463f6a738..f483aa032 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py @@ -15,7 +15,7 @@ from ..kv_transporter import create_kv_transporter from lightllm.utils.error_utils import log_exception from lightllm.utils.envs_utils import get_unique_server_name -from lightllm.utils.process_check import start_parent_check_thread +from lightllm.utils.process_check import install_fatal_thread_excepthook, start_parent_check_thread logger = init_logger(__name__) @@ -41,6 +41,7 @@ def _init_env( task_in_queue: mp.Queue, task_out_queue: mp.Queue, ): + install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ @@ -329,48 +330,76 @@ def update_task_status_loop( time.sleep(0.001) continue - with self.waiting_dict_lock: - tasks = list(self.waiting_dict.values()) - for trans_task in tasks: - if trans_task.xfer_handle is None: - continue + self._update_task_status_once() + time.sleep(0.001) - # 传输任务状态检查 + def _update_task_status_once(self): + with self.waiting_dict_lock: + tasks = list(self.waiting_dict.values()) + for trans_task in tasks: + if trans_task.xfer_handle is None: + continue + + # A single native NIXL error must fail and release this task without + # silently terminating the status thread for every later transfer. + try: ret = self.transporter.check_task_status(trans_task=trans_task) - if ret == "DONE": - trans_task = self.waiting_dict.pop(trans_task.get_key(), None) - if self.transporter.capture_telemetry: - telem = self.transporter.get_xfer_telemetry(trans_task.xfer_handle) + except BaseException as e: + failed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if failed_task is not None: + logger.exception(f"check transfer status failed: {failed_task.to_str()}") + failed_task.error_info = f"check transfer status failed: {str(e)}" + self.failed_queue.put(failed_task) + continue + + if ret == "DONE": + completed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if completed_task is None: + continue + if self.transporter.capture_telemetry: + try: + telem = self.transporter.get_xfer_telemetry(completed_task.xfer_handle) total_us = telem.xferDuration post_us = telem.postDuration backend_us = telem.xferDuration - telem.postDuration - nixl_backend = self.transporter.query_xfer_backend(trans_task.xfer_handle) + nixl_backend = self.transporter.query_xfer_backend(completed_task.xfer_handle) logger.info( - f"write trans task request_id={trans_task.request_id} " - f"kv=[{trans_task.start_kv_index},{trans_task.end_kv_index}) " - f"src_page={trans_task.src_page_index} dst_page={trans_task.dst_page_index} " + f"write trans task request_id={completed_task.request_id} " + f"kv=[{completed_task.start_kv_index},{completed_task.end_kv_index}) " + f"src_page={completed_task.src_page_index} " + f"dst_page={completed_task.dst_page_index} " f"xfer time: {total_us:.3f} us, " f"post time: {post_us:.3f} us, backend time: {backend_us:.3f} us, " f"nixl_backend: {nixl_backend}, total_bytes: {telem.totalBytes}" ) - self.transporter.send_write_done_task_to_decode_node(trans_task) - logger.info( - f"send WRITE done nixl notify " - f"request_id={trans_task.request_id} " - f"kv=[{trans_task.start_kv_index},{trans_task.end_kv_index}) " - f"src_page={trans_task.src_page_index} dst_page={trans_task.dst_page_index}" - ) - self.success_queue.put(trans_task) - elif ret == "ERR": - trans_task = self.waiting_dict.pop(trans_task.get_key(), None) - trans_task.error_info = "xfer error" - self.failed_queue.put(trans_task) - elif trans_task.time_out(): - trans_task = self.waiting_dict.pop(trans_task.get_key(), None) - trans_task.error_info = "time out in update_task_status_loop" - self.failed_queue.put(trans_task) + except BaseException: + logger.exception(f"get transfer telemetry failed: {completed_task.to_str()}") + + try: + self.transporter.send_write_done_task_to_decode_node(completed_task) + except BaseException as e: + logger.exception(f"send WRITE done nixl notify failed: {completed_task.to_str()}") + completed_task.error_info = f"send WRITE done nixl notify failed: {str(e)}" + self.failed_queue.put(completed_task) + continue - time.sleep(0.001) + logger.info( + f"send WRITE done nixl notify " + f"request_id={completed_task.request_id} " + f"kv=[{completed_task.start_kv_index},{completed_task.end_kv_index}) " + f"src_page={completed_task.src_page_index} dst_page={completed_task.dst_page_index}" + ) + self.success_queue.put(completed_task) + elif ret == "ERR": + failed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if failed_task is not None: + failed_task.error_info = "xfer error" + self.failed_queue.put(failed_task) + elif trans_task.time_out(): + failed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if failed_task is not None: + failed_task.error_info = "time out in update_task_status_loop" + self.failed_queue.put(failed_task) @log_exception def success_loop(self): diff --git a/lightllm/utils/process_check.py b/lightllm/utils/process_check.py index 00cc258bf..db508c8a6 100644 --- a/lightllm/utils/process_check.py +++ b/lightllm/utils/process_check.py @@ -8,6 +8,21 @@ logger = init_logger(__name__) +def install_fatal_thread_excepthook(): + """Terminate the process when an unexpected daemon-thread exception escapes.""" + + def exit_process(args): + try: + logger.error( + f"fatal background thread failure in {getattr(args.thread, 'name', 'unknown')}", + exc_info=(args.exc_type, args.exc_value, args.exc_traceback), + ) + finally: + os._exit(1) + + threading.excepthook = exit_process + + def is_process_active(pid): try: process = psutil.Process(pid) diff --git a/unit_tests/test_nixl_kv_transporter_thread_safety.py b/unit_tests/test_nixl_kv_transporter_thread_safety.py index f32ab8852..97276e60c 100644 --- a/unit_tests/test_nixl_kv_transporter_thread_safety.py +++ b/unit_tests/test_nixl_kv_transporter_thread_safety.py @@ -1,10 +1,15 @@ import threading import time +import queue +from types import SimpleNamespace import pytest from lightllm.server.pd_io_struct import PDChunckedTransTask, PDAgentMetadata from lightllm.server.router.model_infer.mode_backend.pd import nixl_kv_transporter +from lightllm.server.router.model_infer.mode_backend.pd.prefill_node_impl.prefill_trans_process import ( + _PrefillTransModule, +) class _FakeBuffer: @@ -247,3 +252,70 @@ def test_active_transfer_defers_removal_and_stale_generation_cannot_remove_recon assert transporter._peer_generations == {"decode": second_generation} assert "decode" in transporter.remote_agents + + +def _status_module(transporter, task): + module = _PrefillTransModule.__new__(_PrefillTransModule) + module.transporter = transporter + module.waiting_dict_lock = threading.Lock() + module.waiting_dict = {task.get_key(): task} + module.success_queue = queue.Queue() + module.failed_queue = queue.Queue() + return module + + +def _status_task(): + return SimpleNamespace( + xfer_handle=object(), + error_info=None, + request_id=1, + start_kv_index=0, + end_kv_index=1, + src_page_index=0, + dst_page_index=0, + get_key=lambda: "task-1", + to_str=lambda: "task-1", + time_out=lambda: False, + ) + + +def test_status_query_failure_routes_task_to_failure_cleanup(): + class FailingTransporter: + capture_telemetry = False + + @staticmethod + def check_task_status(trans_task): + raise RuntimeError("native status failure") + + task = _status_task() + module = _status_module(FailingTransporter(), task) + + module._update_task_status_once() + + assert module.waiting_dict == {} + assert module.failed_queue.get_nowait() is task + assert task.error_info == "check transfer status failed: native status failure" + assert module.success_queue.empty() + + +def test_done_notification_failure_routes_task_to_failure_cleanup(): + class FailingTransporter: + capture_telemetry = False + + @staticmethod + def check_task_status(trans_task): + return "DONE" + + @staticmethod + def send_write_done_task_to_decode_node(trans_task): + raise RuntimeError("notify failure") + + task = _status_task() + module = _status_module(FailingTransporter(), task) + + module._update_task_status_once() + + assert module.waiting_dict == {} + assert module.failed_queue.get_nowait() is task + assert task.error_info == "send WRITE done nixl notify failed: notify failure" + assert module.success_queue.empty() From a4172bed383e902f2fa26508f4b19ff7ee288d5b Mon Sep 17 00:00:00 2001 From: sufubao Date: Thu, 20 Aug 2026 23:14:31 +0800 Subject: [PATCH 21/21] style(pd): apply repository pre-commit formatting Apply the Black 21.12b0 line-width-120 output required by the upstream pre-commit workflow. This commit is mechanical and contains no runtime changes. --- lightllm/server/httpserver/manager.py | 12 +++--------- lightllm/server/httpserver/pd_loop.py | 14 +++----------- .../server/httpserver/test_pd_connection_tasks.py | 4 +--- .../server/httpserver/test_pd_pending_abort.py | 4 +--- 4 files changed, 8 insertions(+), 26 deletions(-) diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index edd2b192b..7e04ba81c 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -311,9 +311,7 @@ def alloc_req_id(self, sampling_params): assert False, "dead code path" return group_request_id - async def _alloc_req_objs( - self, group_request_id, prompt_ids, sampling_params - ) -> List[Req]: + async def _alloc_req_objs(self, group_request_id, prompt_ids, sampling_params) -> List[Req]: """Allocate and initialize request slots without leaking partial allocations.""" alloced_req_indexes = [] req_objs = [] @@ -330,9 +328,7 @@ async def _alloc_req_objs( alloced_req_indexes.append(alloc_req_index) for i, req_index in enumerate(alloced_req_indexes): - req_obj = await self.shm_req_manager.async_get_req_obj_by_index( - req_index - ) + req_obj = await self.shm_req_manager.async_get_req_obj_by_index(req_index) req_objs.append(req_obj) req_obj.init( group_request_id + i, @@ -452,9 +448,7 @@ async def generate( raise PDPrefillNodeStopGenToken(group_request_id=group_request_id) # 申请资源并存储 - req_objs = await self._alloc_req_objs( - group_request_id, prompt_ids, sampling_params - ) + req_objs = await self._alloc_req_objs(group_request_id, prompt_ids, sampling_params) self._log_stage_timing( group_request_id, start_time, diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index 697bb65e1..ccd45181e 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -40,11 +40,7 @@ async def _abort_pd_request( generation_tasks: Dict[int, asyncio.Task], ) -> bool: generation_task = generation_tasks.get(group_req_id) - if ( - generation_task is not None - and not generation_task.done() - and not generation_task.cancelling() - ): + if generation_task is not None and not generation_task.done() and not generation_task.cancelling(): # The request may still be waiting for a shared-memory slot and therefore # not yet be visible to HttpServerManager.abort(). Cancelling its owning # task makes the abort effective at every point in the admission path. @@ -299,14 +295,10 @@ async def _pd_process_generate( except PDPrefillNodeStopGenToken as e: logger.info(f"pd prefill node stop gen token for group_request_id {e.group_request_id}") except asyncio.CancelledError: - logger.info( - f"pd request task cancelled for group_request_id {sampling_params.group_request_id}" - ) + logger.info(f"pd request task cancelled for group_request_id {sampling_params.group_request_id}") raise except BaseException: - logger.exception( - f"PD request generation failed for group_request_id {sampling_params.group_request_id}" - ) + logger.exception(f"PD request generation failed for group_request_id {sampling_params.group_request_id}") raise diff --git a/unit_tests/server/httpserver/test_pd_connection_tasks.py b/unit_tests/server/httpserver/test_pd_connection_tasks.py index 6a4fd8367..eaf2a8e08 100644 --- a/unit_tests/server/httpserver/test_pd_connection_tasks.py +++ b/unit_tests/server/httpserver/test_pd_connection_tasks.py @@ -26,9 +26,7 @@ async def recv(self): websocket = BlockingWebsocket() failure = asyncio.get_running_loop().create_future() - receive_task = asyncio.create_task( - _recv_or_raise_on_background_failure(websocket, (failure,)) - ) + receive_task = asyncio.create_task(_recv_or_raise_on_background_failure(websocket, (failure,))) await recv_started.wait() failure.set_exception(RuntimeError("generation failed")) diff --git a/unit_tests/server/httpserver/test_pd_pending_abort.py b/unit_tests/server/httpserver/test_pd_pending_abort.py index 098744ac3..0ece4d4d8 100644 --- a/unit_tests/server/httpserver/test_pd_pending_abort.py +++ b/unit_tests/server/httpserver/test_pd_pending_abort.py @@ -57,9 +57,7 @@ async def alloc_req_index(): manager.shm_req_manager.async_alloc_req_index = alloc_req_index sampling_params = MagicMock(n=2) - allocation_task = asyncio.create_task( - manager._alloc_req_objs(123, [1, 2], sampling_params) - ) + allocation_task = asyncio.create_task(manager._alloc_req_objs(123, [1, 2], sampling_params)) await waiting_for_second_slot.wait() allocation_task.cancel()