Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 12 additions & 8 deletions lightllm/server/httpserver_for_pd_master/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -211,19 +211,22 @@ async def _generate_one(
block_group_request_id = origin_request_id
p_node = None
d_node = None
prefill_load_released = False
# 只有真正累加过 dispatched_prompt_chars 才需要(且只能)扣减一次。
prefill_load_charged = False

try:
p_node, d_node = await self.select_p_d_node(prompt, origin_sampling_params, multimodal_params)
# 记录当前 P 节点的在途 prompt 负载,首 token 生成后即可释放。
p_node.dispatched_prompt_chars += len(prompt)

history_gen_token_strs = []

if not p_node or not d_node:
logger.error(f"{origin_request_id}: No p_node or d_node found")
raise Exception(f"{origin_request_id}: No p_node or d_node found")

# 记录当前 P 节点的在途 prompt 负载,首 token 生成后即可释放。
p_node.dispatched_prompt_chars += len(prompt)
prefill_load_charged = True

history_gen_token_strs = []

origin_prompt_cache_len = None

for iter_index, block_max_new_tokens in enumerate(max_new_tokens_list):
Expand Down Expand Up @@ -255,9 +258,9 @@ async def _generate_one(
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
if not prefill_load_released:
if prefill_load_charged:
p_node.dispatched_prompt_chars = max(0, p_node.dispatched_prompt_chars - len(prompt))
prefill_load_released = True
prefill_load_charged = False
yield origin_request_id, request_output, metadata, finish_status

await self.remove_req(group_request_id=block_group_request_id)
Expand All @@ -277,8 +280,9 @@ async def _generate_one(
raise e

finally:
if p_node is not None and not prefill_load_released:
if prefill_load_charged:
p_node.dispatched_prompt_chars = max(0, p_node.dispatched_prompt_chars - len(prompt))
prefill_load_charged = False
await self.remove_req(block_group_request_id)
return

Expand Down