Fix KeyError when logprobs=false in completions endpoint (#16095)
This commit is contained in:
@@ -242,9 +242,9 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
output_token_logprobs=content["meta_info"][
|
output_token_logprobs=content["meta_info"][
|
||||||
"output_token_logprobs"
|
"output_token_logprobs"
|
||||||
][n_prev_token:],
|
][n_prev_token:],
|
||||||
output_top_logprobs=content["meta_info"]["output_top_logprobs"][
|
output_top_logprobs=content["meta_info"].get(
|
||||||
n_prev_token:
|
"output_top_logprobs", []
|
||||||
],
|
)[n_prev_token:],
|
||||||
)
|
)
|
||||||
n_prev_tokens[index] = len(
|
n_prev_tokens[index] = len(
|
||||||
content["meta_info"]["output_token_logprobs"]
|
content["meta_info"]["output_token_logprobs"]
|
||||||
@@ -398,10 +398,12 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
logprobs = to_openai_style_logprobs(
|
logprobs = to_openai_style_logprobs(
|
||||||
input_token_logprobs=input_token_logprobs,
|
input_token_logprobs=input_token_logprobs,
|
||||||
input_top_logprobs=input_top_logprobs,
|
input_top_logprobs=input_top_logprobs,
|
||||||
output_token_logprobs=ret_item["meta_info"][
|
output_token_logprobs=ret_item["meta_info"].get(
|
||||||
"output_token_logprobs"
|
"output_token_logprobs", []
|
||||||
],
|
),
|
||||||
output_top_logprobs=ret_item["meta_info"]["output_top_logprobs"],
|
output_top_logprobs=ret_item["meta_info"].get(
|
||||||
|
"output_top_logprobs", []
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Handle hidden states
|
# Handle hidden states
|
||||||
|
|||||||
@@ -156,6 +156,31 @@ class ServingCompletionTestCase(unittest.TestCase):
|
|||||||
# (but might have json_schema from the legacy json_schema field)
|
# (but might have json_schema from the legacy json_schema field)
|
||||||
self.assertIsNone(sampling_params.get("structural_tag"))
|
self.assertIsNone(sampling_params.get("structural_tag"))
|
||||||
|
|
||||||
|
def test_logprobs_false_non_streaming(self):
|
||||||
|
"""Test that logprobs=False doesn't cause KeyError in non-streaming response."""
|
||||||
|
req = CompletionRequest(
|
||||||
|
model="x", prompt="Hello", max_tokens=10, logprobs=False
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_ret = [
|
||||||
|
{
|
||||||
|
"text": " world",
|
||||||
|
"meta_info": {
|
||||||
|
"id": "test-id",
|
||||||
|
"prompt_tokens": 1,
|
||||||
|
"completion_tokens": 2,
|
||||||
|
"finish_reason": {"type": "stop"},
|
||||||
|
"weight_version": "v1",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
response = self.sc._build_completion_response(req, mock_ret, 1234567890)
|
||||||
|
|
||||||
|
self.assertEqual(len(response.choices), 1)
|
||||||
|
self.assertEqual(response.choices[0].text, " world")
|
||||||
|
self.assertEqual(len(response.choices[0].logprobs.top_logprobs), 0)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main(verbosity=2)
|
unittest.main(verbosity=2)
|
||||||
|
|||||||
Reference in New Issue
Block a user