From 982fd860e54118898509eaa0a3e0302e9002ef77 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 4 Aug 2026 17:22:07 +0800 Subject: [PATCH] fix(pd): report real cached_tokens instead of always 0 In PD mode cached_tokens was always 0 in the response usage and access log. prompt_cache_len is only written into shm_req on the prefill node (in _match_radix_cache); the decode node never sets it, so every decode frame carries 0. api_openai overwrites cached_tokens on every stream chunk, so the final usage takes the last decode frame's 0 and drops the real hit reported by the prefill first frame. Fix: in generate(), capture the first block's prefill-reported prompt_cache_len and stamp it onto every yielded frame (single and multi block). In _wait_to_token_package(), track the max prompt_cache_len across frames for the access log instead of popping the last frame's 0. --- .../httpserver_for_pd_master/manager.py | 8 +- .../test_pd_master_cached_tokens.py | 78 +++++++++++++++++++ 2 files changed, 85 insertions(+), 1 deletion(-) create mode 100644 unit_tests/server/httpserver/test_pd_master_cached_tokens.py diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 1a900e969..ab479ee1d 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -187,6 +187,8 @@ async def _generate( logger.error(f"{origin_group_request_id}: No p_node or d_node found") raise Exception(f"{origin_group_request_id}: No p_node or d_node found") + origin_prompt_cache_len = None + for iter_index, block_max_new_tokens in enumerate(max_new_tokens_list): sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params) block_group_request_id = self.id_gen.generate_id() @@ -213,6 +215,9 @@ async def _generate( history_gen_token_strs.append(request_output) prompt_tokens = min(prompt_tokens, metadata["prompt_tokens"]) metadata["prompt_tokens"] = prompt_tokens + if iter_index == 0 and origin_prompt_cache_len is None: + origin_prompt_cache_len = metadata.get("prompt_cache_len", 0) + metadata["prompt_cache_len"] = origin_prompt_cache_len or 0 yield origin_group_request_id, request_output, metadata, finish_status await self.remove_req(group_request_id=block_group_request_id) @@ -381,6 +386,7 @@ async def _wait_to_token_package( out_token_counter = 0 first_token_cost_ms = float("inf") + prompt_cache_len = 0 group_request_id = sampling_params.group_request_id unfinished_count = sampling_params.best_of is_first_token = True @@ -396,6 +402,7 @@ async def _wait_to_token_package( prompt_tokens = metadata["prompt_tokens"] out_token_counter += 1 + prompt_cache_len = max(prompt_cache_len, metadata.get("prompt_cache_len", 0)) sub_req_id_to_mtp_accepted_token_num[sub_req_id] = metadata.get("mtp_accepted_token_num", 0) if is_first_token: first_token_cost_ms = (time.time() - start_time) * 1000 @@ -414,7 +421,6 @@ async def _wait_to_token_package( self.per_token_costs.add(mean_per_token_cost_time_ms) x_request_id = request.headers.get("X-Request-Id", "") x_session_id = request.headers.get("X-Session-Id", "") - prompt_cache_len = metadata.pop("prompt_cache_len", 0) prompt_cache_ratio = prompt_cache_len / prompt_tokens mtp_avg_token_per_step = out_token_counter / max( (out_token_counter - sum(sub_req_id_to_mtp_accepted_token_num.values())), 1 diff --git a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py new file mode 100644 index 000000000..1e65ec990 --- /dev/null +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -0,0 +1,78 @@ +import asyncio +import copy +from types import SimpleNamespace + +import pytest + +from lightllm.server.core.objs import FinishStatus, SamplingParams +from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster + + +def _make_manager(monkeypatch): + monkeypatch.setattr( + "lightllm.server.httpserver.manager.HttpServerManager._check_and_repair_length", + classmethod(lambda cls, *a, **k: asyncio.sleep(0)), + ) + monkeypatch.setattr(SamplingParams, "from_buffer_copy", classmethod(lambda cls, other: copy.copy(other))) + mgr = object.__new__(HttpServerManagerForPDMaster) + mgr.running_request_count = 0 + counter = [0] + + def gen_id(): + counter[0] += 1 + return counter[0] + + mgr.id_gen = SimpleNamespace(generate_id=gen_id) + mgr.metric_client = SimpleNamespace(counter_inc=lambda *a, **k: None, histogram_observe=lambda *a, **k: None) + mgr.tokens = lambda *a, **k: 10 + mgr._log_req_header = lambda *a, **k: asyncio.sleep(0) + mgr.select_p_d_node = lambda *a, **k: asyncio.sleep(0, result=(1, 1)) + mgr.remove_req = lambda *a, **k: asyncio.sleep(0) + return mgr + + +def _collect(mgr, sampling_params, monkeypatch, split): + mgr._split_max_new_tokens = lambda *a, **k: list(split) + + async def fake_wait(p_node, d_node, start_time, prompt, sp, multimodal_params, request): + sub_req_id = sp.group_request_id + hit = sp.max_new_tokens * 10 + yield sub_req_id, "x", {"prompt_tokens": 10, "prompt_cache_len": hit}, FinishStatus() + for _ in range(2): + yield sub_req_id, "y", {"prompt_tokens": 10, "prompt_cache_len": 0}, FinishStatus() + yield sub_req_id, "z", {"prompt_tokens": 10, "prompt_cache_len": 0}, FinishStatus(FinishStatus.FINISHED_STOP) + + monkeypatch.setattr(mgr, "_wait_to_token_package", fake_wait) + + async def run(): + out = [] + async for sub_id, out_str, metadata, finish in mgr.generate( + prompt="hello", + sampling_params=sampling_params, + multimodal_params=SimpleNamespace(images=[], audios=[], verify_and_preload=lambda req: asyncio.sleep(0)), + request=None, + ): + out.append(metadata.get("prompt_cache_len", -1)) + return out + + return asyncio.run(run()) + + +def test_single_block_prefill_hit_persists_past_decode_zeros(monkeypatch): + mgr = _make_manager(monkeypatch) + sp = SamplingParams() + sp.max_new_tokens = 3 + sp.best_of = 1 + sp.group_request_id = 0 + cached = _collect(mgr, sp, monkeypatch, split=[3]) + assert cached and all(c == 30 for c in cached), cached + + +def test_multi_block_keeps_first_block_hit(monkeypatch): + mgr = _make_manager(monkeypatch) + sp = SamplingParams() + sp.max_new_tokens = 5 + sp.best_of = 1 + sp.group_request_id = 0 + cached = _collect(mgr, sp, monkeypatch, split=[3, 2]) + assert cached[-1] == 30, cached