Skip to content
Merged
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
23 changes: 16 additions & 7 deletions megatron/core/inference/engines/dynamic_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -1218,12 +1218,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
# [<prompt log probs...>, <sampled token log prob>], 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
Expand Down Expand Up @@ -1251,7 +1260,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,
Expand All @@ -1264,7 +1273,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))

Expand Down
63 changes: 63 additions & 0 deletions tests/unit_tests/inference/engines/test_dynamic_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Expand Down
Loading