diff --git a/lmdeploy/serve/openai/chat_completions/logprobs.py b/lmdeploy/serve/openai/chat_completions/logprobs.py index ea0d5b464f..236c75ada0 100644 --- a/lmdeploy/serve/openai/chat_completions/logprobs.py +++ b/lmdeploy/serve/openai/chat_completions/logprobs.py @@ -13,7 +13,8 @@ def _create_chat_completion_logprobs(tokenizer: PreTrainedTokenizerBase, token_ids: list[int] | None = None, logprobs: list[dict[int, float]] - | None = None): + | None = None, + top_logprobs: int = 0): """Create openai LogProbs for chat.completion. Args: @@ -21,6 +22,8 @@ def _create_chat_completion_logprobs(tokenizer: PreTrainedTokenizerBase, token_ids (list[int]): output token ids. logprobs (list[dict[int, float]]): the top logprobs for each output position. + top_logprobs (int): the number of most likely tokens to return at + each output position. Returns: ChoiceLogprobs: logprob result. """ @@ -33,7 +36,15 @@ def _create_chat_completion_logprobs(tokenizer: PreTrainedTokenizerBase, bytes=[], logprob=0.0, top_logprobs=[]) - for top_id, prob in tops.items(): + top_ids = sorted(tops, key=tops.get, reverse=True) + if len(top_ids) > top_logprobs: + # Drop the extra selected row the engine appends when the + # selected token is outside the model top-k. + top_ids = [top_id for top_id in top_ids if top_id != token_id] + top_ids = top_ids[:top_logprobs] + for top_id, prob in sorted(tops.items(), key=lambda x: x[1], reverse=True): + if top_id != token_id and top_id not in top_ids: + continue token = tokenizer.convert_ids_to_tokens(top_id) if isinstance(token, bytes): _bytes = list(token) @@ -44,7 +55,7 @@ def _create_chat_completion_logprobs(tokenizer: PreTrainedTokenizerBase, item.token = token item.bytes = _bytes item.logprob = prob - else: + if top_id in top_ids: item.top_logprobs.append( TopLogprob(token=token, bytes=_bytes, logprob=prob)) content.append(item) diff --git a/lmdeploy/serve/openai/chat_completions/serving.py b/lmdeploy/serve/openai/chat_completions/serving.py index ea57c43221..a146ce8099 100644 --- a/lmdeploy/serve/openai/chat_completions/serving.py +++ b/lmdeploy/serve/openai/chat_completions/serving.py @@ -274,7 +274,7 @@ async def _completion_stream_generator() -> AsyncGenerator[str, None]: output_token_logprobs = None if request.logprobs and chunk.logprobs: logprobs = _create_chat_completion_logprobs( - tokenizer, chunk.token_ids, chunk.logprobs) + tokenizer, chunk.token_ids, chunk.logprobs, request.top_logprobs or 0) if request.return_logprob and chunk.logprobs: output_token_logprobs = _create_output_token_logprobs( chunk.token_ids, chunk.logprobs) @@ -344,7 +344,7 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: logprobs = None if request.logprobs and len(res.logprobs): logprobs = _create_chat_completion_logprobs( - tokenizer, res.token_ids, res.logprobs) + tokenizer, res.token_ids, res.logprobs, request.top_logprobs or 0) output_token_logprobs = None if request.return_logprob and len(res.logprobs): output_token_logprobs = _create_output_token_logprobs( diff --git a/tests/test_lmdeploy/serve/openai/chat_completions/test_logprobs.py b/tests/test_lmdeploy/serve/openai/chat_completions/test_logprobs.py new file mode 100644 index 0000000000..43bfe8bf89 --- /dev/null +++ b/tests/test_lmdeploy/serve/openai/chat_completions/test_logprobs.py @@ -0,0 +1,46 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import pytest + +from lmdeploy.serve.openai.chat_completions.logprobs import _create_chat_completion_logprobs + + +class _Tokenizer: + + def convert_ids_to_tokens(self, token_id): + return f'tok{token_id}' + + +def _top(item): + return [(top.token, top.logprob) for top in item.top_logprobs] + + +@pytest.mark.parametrize('top_logprobs', [1, 3]) +def test_top_logprobs_include_selected_token_in_top_k(top_logprobs): + # Engines return the model top-k; the selected token 1 is the most likely. + tops = {1: -0.1, 2: -1.0, 3: -2.0} + tops = {k: v for k, v in list(tops.items())[:top_logprobs]} + result = _create_chat_completion_logprobs(_Tokenizer(), [1], [tops], top_logprobs) + + item = result.content[0] + assert (item.token, item.logprob) == ('tok1', -0.1) + assert _top(item) == [(f'tok{k}', v) for k, v in tops.items()] + + +def test_top_logprobs_drop_selected_token_outside_top_k(): + # The selected token 4 is appended after the model top-2. + tops = {4: -3.0, 1: -0.1, 2: -1.0} + result = _create_chat_completion_logprobs(_Tokenizer(), [4], [tops], 2) + + item = result.content[0] + assert (item.token, item.logprob) == ('tok4', -3.0) + assert _top(item) == [('tok1', -0.1), ('tok2', -1.0)] + + +def test_top_logprobs_zero_returns_empty_list(): + # logprobs=True without top_logprobs still asks the engine for one entry. + tops = {4: -3.0, 1: -0.1} + result = _create_chat_completion_logprobs(_Tokenizer(), [4], [tops], 0) + + item = result.content[0] + assert (item.token, item.logprob) == ('tok4', -3.0) + assert item.top_logprobs == []