From 8a2929969cd28a2d5e3ef479176dedf9d3010b18 Mon Sep 17 00:00:00 2001 From: Keshav Santhanam Date: Thu, 4 Jun 2026 11:18:03 -0700 Subject: [PATCH 1/2] Fix prompt logprobs Signed-off-by: Keshav Santhanam --- .../core/inference/engines/dynamic_engine.py | 23 ++++--- .../inference/engines/test_dynamic_engine.py | 63 +++++++++++++++++++ 2 files changed, 79 insertions(+), 7 deletions(-) diff --git a/megatron/core/inference/engines/dynamic_engine.py b/megatron/core/inference/engines/dynamic_engine.py index 8a43eb0f7ae..b3a81199dd5 100644 --- a/megatron/core/inference/engines/dynamic_engine.py +++ b/megatron/core/inference/engines/dynamic_engine.py @@ -1211,12 +1211,21 @@ def post_process_requests( keep = request.sampling_params.num_tokens_to_generate - len( request.generated_tokens ) + num_tokens_before_trim = len(tokens) tokens = tokens[:keep] - # Trim log probs / top-n to match so the counts stay in sync. - if request_log_probs is not None: - request_log_probs = request_log_probs[:keep] - if top_n_logprobs is not None and req_idx in top_n_logprobs: - top_n_logprobs[req_idx] = top_n_logprobs[req_idx][:keep] + # Drop only the excess *trailing* log probs / top-n so the counts stay + # in sync. We must trim from the end, not the front: on a prefill step + # request_log_probs covers the whole prompt and is laid out as + # [, ], so front-slicing + # (e.g. [:keep] with keep == 0 when num_tokens_to_generate == 0) would + # discard the prompt log probs that echo+logprobs requests need. In a + # decode step all entries are generated, so trailing == front-equivalent. + num_dropped = num_tokens_before_trim - len(tokens) + if num_dropped > 0: + if request_log_probs is not None: + request_log_probs = request_log_probs[:-num_dropped] + if top_n_logprobs is not None and req_idx in top_n_logprobs: + top_n_logprobs[req_idx] = top_n_logprobs[req_idx][:-num_dropped] if request_id not in self.stop_word_being_finished_ids: is_first_token = len(request.generated_tokens) == 0 request.generated_tokens += tokens @@ -1244,7 +1253,7 @@ def post_process_requests( ) if first_token_event is None: first_token_event = event - if is_first_token: + if is_first_token and tokens: if not self.track_generated_token_events: first_token_event = DynamicInferenceEvent( type=DynamicInferenceEventType.GENERATED_TOKEN, @@ -1257,7 +1266,7 @@ def post_process_requests( # non-logging steps (async_forward skips the event sync), # so gate the update to keep the metric a truthful sparse # sample instead of polluting it with zeros. - if step_time > 0: + if step_time > 0 and tokens: per_token_step_time = step_time / len(tokens) request.tpot.extend([per_token_step_time] * len(tokens)) diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine.py b/tests/unit_tests/inference/engines/test_dynamic_engine.py index 2f03c8cb7aa..351d26bdebb 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine.py @@ -1151,6 +1151,69 @@ def test_return_log_probs(self): f"log_prob {log_prob} is out of expected range [-50.0, 0.0]" ) + @pytest.mark.skipif( + not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + ) + @torch.inference_mode() + def test_return_prompt_log_probs_with_zero_tokens_to_generate(self): + """Prompt log probs must be returned when scoring only (num_tokens_to_generate=0). + + Regression test for a prefill-step trimming bug: when a request generates + no tokens, the end-of-generation trim set ``keep=0`` and front-sliced + ``request_log_probs[:0]``, discarding every prompt log prob (in a prefill + step ``request_log_probs`` covers the whole prompt, with the disposable + sampled-token log prob at the tail). The fix trims the excess *trailing* + log probs instead. This is the path exercised by loglikelihood / echo + evaluations (e.g. lm-eval-harness sends ``max_tokens=0``). + """ + env = self._run_test( + return_log_probs=True, + materialize_only_last_token_logits=False, + skip_prompt_log_probs=False, + num_tokens_to_generate=0, + ) + + validated_any = False + for request in env.requests: + if request.status != Status.COMPLETED: + continue + + # No tokens were requested, so none should be generated. + assert len(request.generated_tokens) == 0, ( + f"Request {request.request_id}: expected 0 generated tokens, " + f"got {len(request.generated_tokens)}" + ) + assert request.generated_log_probs is None or len(request.generated_log_probs) == 0, ( + f"Request {request.request_id}: expected no generated log probs, got " + f"{len(request.generated_log_probs) if request.generated_log_probs else 0}" + ) + + # The full set of prompt log probs (all tokens except the first) must + # still be present -- before the fix this list was empty. + prompt_len = len(request.prompt_tokens) + assert request.prompt_log_probs is not None, ( + f"Request {request.request_id}: prompt_log_probs should not be None " + f"when scoring with num_tokens_to_generate=0" + ) + assert len(request.prompt_log_probs) == prompt_len - 1, ( + f"Request {request.request_id}: Expected {prompt_len - 1} prompt log probs, " + f"got {len(request.prompt_log_probs)}" + ) + for i, log_prob in enumerate(request.prompt_log_probs): + assert log_prob is not None, ( + f"Request {request.request_id}, prompt token {i}: log_prob is None" + ) + assert not math.isnan(log_prob) and not math.isinf(log_prob), ( + f"Request {request.request_id}, prompt token {i}: log_prob {log_prob} invalid" + ) + assert -50.0 <= log_prob <= 0.0, ( + f"Request {request.request_id}, prompt token {i}: " + f"log_prob {log_prob} is out of expected range [-50.0, 0.0]" + ) + validated_any = True + + assert validated_any, "No completed requests were validated" + @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) From 542e8226118b43e76b02ec9ba70b5911a4cf0a5a Mon Sep 17 00:00:00 2001 From: Keshav Santhanam Date: Thu, 4 Jun 2026 16:04:50 -0700 Subject: [PATCH 2/2] Linting Signed-off-by: Keshav Santhanam --- .../inference/engines/test_dynamic_engine.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine.py b/tests/unit_tests/inference/engines/test_dynamic_engine.py index 151395e5b98..038750a22b6 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine.py @@ -1200,12 +1200,12 @@ def test_return_prompt_log_probs_with_zero_tokens_to_generate(self): f"got {len(request.prompt_log_probs)}" ) for i, log_prob in enumerate(request.prompt_log_probs): - assert log_prob is not None, ( - f"Request {request.request_id}, prompt token {i}: log_prob is None" - ) - assert not math.isnan(log_prob) and not math.isinf(log_prob), ( - f"Request {request.request_id}, prompt token {i}: log_prob {log_prob} invalid" - ) + assert ( + log_prob is not None + ), f"Request {request.request_id}, prompt token {i}: log_prob is None" + assert not math.isnan(log_prob) and not math.isinf( + log_prob + ), f"Request {request.request_id}, prompt token {i}: log_prob {log_prob} invalid" assert -50.0 <= log_prob <= 0.0, ( f"Request {request.request_id}, prompt token {i}: " f"log_prob {log_prob} is out of expected range [-50.0, 0.0]"