Skip to content
Open
Show file tree
Hide file tree
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
5 changes: 5 additions & 0 deletions lightllm/server/api_openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -349,6 +349,9 @@ async def chat_completions_impl(request: ChatCompletionRequest, raw_request: Req

sampling_params = SamplingParams()
sampling_params.init(tokenizer=g_objs.httpserver_manager.tokenizer, **sampling_params_dict)
# The chat API does not expose output-token logprobs. PD nodes use this
# transport-only marker to avoid forwarding unused per-token metadata.
sampling_params.return_output_logprobs = False

sampling_params.verify()
results_generator = g_objs.httpserver_manager.generate(
Expand Down Expand Up @@ -843,6 +846,7 @@ async def completions_impl(request: CompletionRequest, raw_request: Request) ->

sampling_params = SamplingParams()
sampling_params.init(tokenizer=g_objs.httpserver_manager.tokenizer, **sampling_params_dict)
sampling_params.return_output_logprobs = request.logprobs is not None
sampling_params.verify()

# v1/completions does not support multimodal inputs, so we use an empty MultimodalParams
Expand Down Expand Up @@ -880,6 +884,7 @@ async def process_single_prompt(prompt: Union[str, List[int]]):
if len(prompts) > 1:
individual_sampling_params = SamplingParams()
individual_sampling_params.init(tokenizer=g_objs.httpserver_manager.tokenizer, **sampling_params_dict)
individual_sampling_params.return_output_logprobs = request.logprobs is not None
individual_sampling_params.verify()
else:
individual_sampling_params = sampling_params
Expand Down
4 changes: 4 additions & 0 deletions lightllm/server/httpserver/pd_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,6 +234,7 @@ async def _pd_process_generate(
pd_upload_websocket: ClientConnection,
pd_event: asyncio.Event,
):
return_output_logprobs = getattr(sampling_params, "return_output_logprobs", True)
try:
async for sub_req_id, request_output, metadata, finish_status in manager.generate(
prompt=prompt,
Expand All @@ -244,6 +245,9 @@ async def _pd_process_generate(
pd_event=pd_event,
):
metadata["node_mode"] = manager.args.run_mode
if not return_output_logprobs:
for key in ("id", "logprob", "cumlogprob", "special", "logprobs"):
metadata.pop(key, None)
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}")
Expand Down
4 changes: 4 additions & 0 deletions lightllm/server/httpserver_for_pd_master/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,7 +160,9 @@ async def _generate(
self, prompt_ids=fake_prompt_ids, sampling_params=sampling_params
)

return_output_logprobs = getattr(sampling_params, "return_output_logprobs", True)
origin_sampling_params = SamplingParams.from_buffer_copy(sampling_params)
origin_sampling_params.return_output_logprobs = return_output_logprobs
origin_group_request_id = self.id_gen.generate_id()

# Record one user request even when it is expanded into multiple independent
Expand All @@ -174,6 +176,7 @@ async def _generate(
generators = []
for choice_index in range(choice_count):
choice_sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params)
choice_sampling_params.return_output_logprobs = return_output_logprobs
choice_sampling_params.n = 1
choice_sampling_params.best_of = 1
generators.append(
Expand Down Expand Up @@ -224,6 +227,7 @@ async def _generate_one(

for iter_index, block_max_new_tokens in enumerate(max_new_tokens_list):
sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params)
sampling_params.return_output_logprobs = getattr(origin_sampling_params, "return_output_logprobs", True)
block_group_request_id = self.id_gen.generate_id()
sampling_params.group_request_id = block_group_request_id
logger.info(f"pd log gen sub req id {block_group_request_id} for main req id {origin_request_id}")
Expand Down
Loading