diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 1a900e969..8d621eaca 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -339,6 +339,9 @@ async def fetch_pd_stream( ) first_token_gen = False + first_token_package = None + pending_token_list = [] + prefill_prompt_cache_len = None while True: await req_status.wait_to_ready() if await request.is_disconnected(): @@ -352,17 +355,33 @@ async def fetch_pd_stream( output_index = metadata.get("count_output_tokens") # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 if output_index == 1: + node_run_mode = metadata.pop("node_mode", None) + if node_run_mode == "prefill": + prefill_prompt_cache_len = metadata.get("prompt_cache_len", 0) if first_token_gen is False: first_token_gen = True - node_run_mode = 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 + first_token_package = (sub_req_id, request_output, metadata, finish_status) else: - continue + if first_token_package is None: + yield sub_req_id, request_output, metadata, finish_status + else: + pending_token_list.append((sub_req_id, request_output, metadata, finish_status)) else: - yield sub_req_id, request_output, metadata, finish_status + if first_token_package is None: + yield sub_req_id, request_output, metadata, finish_status + else: + pending_token_list.append((sub_req_id, request_output, metadata, finish_status)) + + if first_token_package is not None and prefill_prompt_cache_len is not None: + first_token_package[2]["prompt_cache_len"] = prefill_prompt_cache_len + ready_token_list = [first_token_package, *pending_token_list] + first_token_package = None + pending_token_list.clear() + for ready_token in ready_token_list: + yield ready_token return diff --git a/unit_tests/server/httpserver/test_pd_master_token_race.py b/unit_tests/server/httpserver/test_pd_master_token_race.py new file mode 100644 index 000000000..fc3ed60fc --- /dev/null +++ b/unit_tests/server/httpserver/test_pd_master_token_race.py @@ -0,0 +1,84 @@ +import asyncio +import pickle +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 + + +class _FakeReqStatus: + def __init__(self, req_id, p_node, d_node): + self.prefill_prompt_ids_event = SimpleNamespace(prompt_ids=[1, 2, 3]) + self.up_status_event = SimpleNamespace( + upkv_status=SimpleNamespace(pd_kv_trans_params=pickle.dumps(SimpleNamespace())) + ) + self.token_batches = [ + [ + ( + req_id, + "decode-first", + {"count_output_tokens": 1, "node_mode": "decode", "prompt_cache_len": 0}, + FinishStatus(), + ) + ], + [(req_id, "decode-second", {"count_output_tokens": 2, "node_mode": "decode"}, FinishStatus())], + [ + ( + req_id, + "prefill-first", + {"count_output_tokens": 1, "node_mode": "prefill", "prompt_cache_len": 7}, + FinishStatus(FinishStatus.FINISHED_LENGTH), + ) + ], + ] + + async def wait_to_ready(self): + await asyncio.sleep(0) + + async def can_read(self, req_id_to_out_inf): + return bool(self.token_batches) + + async def pop_all_tokens(self): + return self.token_batches.pop(0) + + +class _FakeWebSocket: + async def send_bytes(self, data): + pass + + +def _make_manager(monkeypatch): + monkeypatch.setattr(pd_master_manager, "ReqStatus", _FakeReqStatus) + + async def _noop(*a, **k): + return None + + mgr = object.__new__(HttpServerManagerForPDMaster) + mgr.args = SimpleNamespace(pd_node_id=1) + mgr.req_id_to_out_inf = {} + mgr._wait_for_event_or_disconnect = _noop + return mgr + + +def test_decode_first_token_uses_later_prefill_cache_metadata(monkeypatch): + mgr = _make_manager(monkeypatch) + sampling_params = SamplingParams() + sampling_params.group_request_id = 11 + sampling_params.max_new_tokens = 2 + node = SimpleNamespace(websocket=_FakeWebSocket()) + request = SimpleNamespace(is_disconnected=lambda: asyncio.sleep(0, result=False)) + + async def run(): + stream = mgr.fetch_pd_stream(node, node, "hello", sampling_params, SimpleNamespace(), request) + first = await anext(stream) + second = await anext(stream) + await stream.aclose() + return first, second + + first, second = asyncio.run(run()) + assert first[1] == "decode-first" + assert first[2]["prompt_cache_len"] == 7 + assert second[1] == "decode-second"