fix: preserve output logprobs without input logprobs (#34627)
Signed-off-by: jain-ria <riajain@NVIDIA.com>
This commit is contained in:
@@ -2655,11 +2655,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
recv_obj: BatchStrOutput,
|
recv_obj: BatchStrOutput,
|
||||||
recv_obj_index: int,
|
recv_obj_index: int,
|
||||||
):
|
):
|
||||||
if recv_obj.input_token_logprobs_val is None:
|
|
||||||
return
|
|
||||||
|
|
||||||
if (
|
if (
|
||||||
len(recv_obj.input_token_logprobs_val) > 0
|
recv_obj.input_token_logprobs_val is not None
|
||||||
|
and len(recv_obj.input_token_logprobs_val) > 0
|
||||||
and recv_obj.input_token_logprobs_val[recv_obj_index] is not None
|
and recv_obj.input_token_logprobs_val[recv_obj_index] is not None
|
||||||
):
|
):
|
||||||
state.input_token_logprobs_val.extend(
|
state.input_token_logprobs_val.extend(
|
||||||
@@ -2668,12 +2666,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
state.input_token_logprobs_idx.extend(
|
state.input_token_logprobs_idx.extend(
|
||||||
recv_obj.input_token_logprobs_idx[recv_obj_index]
|
recv_obj.input_token_logprobs_idx[recv_obj_index]
|
||||||
)
|
)
|
||||||
state.output_token_logprobs_val.extend(
|
if (
|
||||||
recv_obj.output_token_logprobs_val[recv_obj_index]
|
recv_obj.output_token_logprobs_val is not None
|
||||||
)
|
and recv_obj.output_token_logprobs_val[recv_obj_index] is not None
|
||||||
state.output_token_logprobs_idx.extend(
|
):
|
||||||
recv_obj.output_token_logprobs_idx[recv_obj_index]
|
state.output_token_logprobs_val.extend(
|
||||||
)
|
recv_obj.output_token_logprobs_val[recv_obj_index]
|
||||||
|
)
|
||||||
|
state.output_token_logprobs_idx.extend(
|
||||||
|
recv_obj.output_token_logprobs_idx[recv_obj_index]
|
||||||
|
)
|
||||||
|
|
||||||
if top_logprobs_num > 0:
|
if top_logprobs_num > 0:
|
||||||
if len(recv_obj.input_top_logprobs_val) > 0:
|
if len(recv_obj.input_top_logprobs_val) > 0:
|
||||||
|
|||||||
@@ -57,6 +57,7 @@ _EXACT_VAL_ROWS = [None, [-0.5, -2.5], [-0.25, -1.5], [-0.125, -4.0]]
|
|||||||
class _TokenizerManagerStub:
|
class _TokenizerManagerStub:
|
||||||
"""Borrow the real logprob meta_info methods without a full manager."""
|
"""Borrow the real logprob meta_info methods without a full manager."""
|
||||||
|
|
||||||
|
convert_logprob_style = TokenizerManager.convert_logprob_style
|
||||||
add_logprob_to_meta_info = TokenizerManager.add_logprob_to_meta_info
|
add_logprob_to_meta_info = TokenizerManager.add_logprob_to_meta_info
|
||||||
detokenize_logprob_tokens = TokenizerManager.detokenize_logprob_tokens
|
detokenize_logprob_tokens = TokenizerManager.detokenize_logprob_tokens
|
||||||
detokenize_top_logprobs_tokens = TokenizerManager.detokenize_top_logprobs_tokens
|
detokenize_top_logprobs_tokens = TokenizerManager.detokenize_top_logprobs_tokens
|
||||||
@@ -543,6 +544,32 @@ class TestMetaInfoFromSchedulerArrays(CustomTestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestTokenizerManagerLogprobs(CustomTestCase):
|
||||||
|
def test_output_logprobs_without_input_logprobs(self):
|
||||||
|
state = _make_state(return_logprob=True, top_logprobs_num=0)
|
||||||
|
recv_obj = SimpleNamespace(
|
||||||
|
input_token_logprobs_val=None,
|
||||||
|
input_token_logprobs_idx=None,
|
||||||
|
output_token_logprobs_val=[[-0.25]],
|
||||||
|
output_token_logprobs_idx=[[42]],
|
||||||
|
)
|
||||||
|
meta_info = {}
|
||||||
|
|
||||||
|
_TokenizerManagerStub().convert_logprob_style(
|
||||||
|
meta_info,
|
||||||
|
state,
|
||||||
|
top_logprobs_num=0,
|
||||||
|
token_ids_logprob=None,
|
||||||
|
return_text_in_logprobs=False,
|
||||||
|
recv_obj=recv_obj,
|
||||||
|
recv_obj_index=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(meta_info["input_token_logprobs"], [])
|
||||||
|
self.assertEqual(meta_info["output_token_logprobs"], [(-0.25, 42, None)])
|
||||||
|
self.assertEqual(meta_info["output_token_logprobs_length"], 1)
|
||||||
|
|
||||||
|
|
||||||
def _make_batch_token_id_output(**overrides) -> BatchTokenIDOutput:
|
def _make_batch_token_id_output(**overrides) -> BatchTokenIDOutput:
|
||||||
"""A two-request BatchTokenIDOutput with the required fields stubbed."""
|
"""A two-request BatchTokenIDOutput with the required fields stubbed."""
|
||||||
n = 2
|
n = 2
|
||||||
|
|||||||
Reference in New Issue
Block a user