From 9e3c13e4fa1c25b9e31d26f8f8495e49757e1f81 Mon Sep 17 00:00:00 2001 From: Thor Johnsen <41591019+thorjohnsen@users.noreply.github.com> Date: Fri, 2 Oct 2026 16:17:41 +0000 Subject: [PATCH] [https://nvbugs/6838020][fix] Tolerate near-tie greedy flips in KV pool rebalance accuracy test test_rebalance_matches_baseline[overlap] failed on B300 because prompt 1 diverged at generated token 10, where " excels" leads " revolutionized" by 0.02 logits in fp32. In bf16 that is one ulp, so any change in accumulation order can flip the greedy pick. Both continuations are fluent text, which points to a near-tie flip rather than KV corruption. B200, GB200 and GB300 passed the overlap variant at the same commit. Request top-2 raw logprobs and compare outputs up to their first divergence. A divergence passes only if, in both arms, the two diverging tokens are that position's top-2 and lie within 0.5 nats of each other. A divergence at a confident position, a token outside the top-2, or one arm stopping early still fails. Signed-off-by: Thor Johnsen <41591019+thorjohnsen@users.noreply.github.com> --- .../test_kv_pool_rebalance_accuracy.py | 94 +++++++++++++++++-- 1 file changed, 84 insertions(+), 10 deletions(-) diff --git a/tests/integration/defs/accuracy/test_kv_pool_rebalance_accuracy.py b/tests/integration/defs/accuracy/test_kv_pool_rebalance_accuracy.py index fd8c592b2308..50abe9c76283 100644 --- a/tests/integration/defs/accuracy/test_kv_pool_rebalance_accuracy.py +++ b/tests/integration/defs/accuracy/test_kv_pool_rebalance_accuracy.py @@ -73,7 +73,22 @@ def _inject_pool_ratio_mismatch(llm: LLM, *, skew: float = 2.0) -> None: "The quick brown fox jumps over the lazy dog. " * 40, ] -_SAMPLING = SamplingParams(max_tokens=64, temperature=0.0, top_k=1) +# Top-2 raw logprobs per generated token, so a divergence between the arms can +# be classified as a near-tie flip or a real accuracy break (see +# _assert_tokens_match_or_near_tie). LogprobMode.RAW (the default) computes +# them from the unprocessed logits, so top_k=1 does not truncate them. +_SAMPLING = SamplingParams(max_tokens=64, temperature=0.0, top_k=1, logprobs=2) + +# Largest top-1/top-2 logprob gap, in nats, at which the two arms may pick +# different greedy tokens. Greedy decode is only exact up to floating-point +# reassociation: anything that changes accumulation order (batch composition, +# kernel selection) moves bf16 logits by a few ulps, which is 0.125 per ulp for +# logits in [16, 32). Candidates that close can swap without anything being +# wrong. nvbugs/6838020 hit exactly that on B300 -- prompt 1 at the token after +# "...architecture that", where " excels" leads " revolutionized" by 0.02 in +# fp32 and by one ulp in bf16. Corrupted KV moves logits far more than this, +# so a divergence at a confident position still fails. +_NEAR_TIE_LOGPROB_MARGIN = 0.5 def _vswa_kv_cache_config(*, enable_rebalance: bool) -> KvCacheConfig: @@ -95,7 +110,10 @@ def _vswa_kv_cache_config(*, enable_rebalance: bool) -> KvCacheConfig: def _generate_tokens(*, model_path: str, disable_overlap: bool, enable_rebalance: bool): - """Run one LLM, return list[list[int]] of generated token ids. + """Run one LLM, return a (token_ids, logprobs) pair per prompt. + + ``logprobs`` holds one ``{token_id: Logprob}`` dict per generated token, + covering that position's top-2 candidates. Note: the ratio-injection helper requires direct access to the in-process PyExecutor, so the test runs in single-process worker @@ -135,15 +153,75 @@ def _generate_tokens(*, model_path: str, disable_overlap: bool, enable_rebalance "supposed to hold pool ratios fixed." ) - return [list(o.outputs[0].token_ids) for o in outputs] + results = [] + for o in outputs: + token_ids = list(o.outputs[0].token_ids) + logprobs = list(o.outputs[0].logprobs) + # The near-tie check indexes logprobs by token position, so a + # missing or misaligned list would make it fail for the wrong reason. + assert len(logprobs) == len(token_ids), ( + f"expected one logprobs entry per generated token, got " + f"{len(logprobs)} for {len(token_ids)} tokens" + ) + results.append((token_ids, logprobs)) + return results + + +def _assert_tokens_match_or_near_tie(prompt_idx: int, baseline, treated) -> None: + """Assert both arms decode identically, apart from at most one near-tie flip. + + Up to their first divergence the outputs must match token for token. The + divergence passes only if, in *both* arms, the two diverging tokens are that + position's top-2 candidates and are within ``_NEAR_TIE_LOGPROB_MARGIN`` of + each other. That means the arms chose opposite sides of a near-tie while + still agreeing on the two leading candidates. Tokens after the divergence + are not compared, because from there each arm continues a different prefix. + """ + b_tokens, b_logprobs = baseline + t_tokens, t_logprobs = treated + context = ( + f"prompt {prompt_idx}: rebalance changed greedy-decode output\n" + f" baseline: {b_tokens[:16]}...\n" + f" treated: {t_tokens[:16]}..." + ) + + k = next((j for j, (b, t) in enumerate(zip(b_tokens, t_tokens)) if b != t), None) + if k is None: + # One output is a prefix of the other, so one arm stopped early. The + # stopping decision does not show up as a token we could look up in + # the logprobs, so the near-tie check cannot vouch for it. + assert len(b_tokens) == len(t_tokens), ( + f"{context}\n outputs agree for {min(len(b_tokens), len(t_tokens))} " + f"tokens, then one arm stops ({len(b_tokens)} vs {len(t_tokens)} tokens)" + ) + return + + b_tok, t_tok = b_tokens[k], t_tokens[k] + for arm, logprobs in (("baseline", b_logprobs), ("treated", t_logprobs)): + top = logprobs[k] + assert b_tok in top and t_tok in top, ( + f"{context}\n diverged at token {k} ({b_tok} vs {t_tok}), but the " + f"{arm} arm's top-2 there is {sorted(top)}: not a near-tie flip" + ) + margin = abs(top[b_tok].logprob - top[t_tok].logprob) + assert margin <= _NEAR_TIE_LOGPROB_MARGIN, ( + f"{context}\n diverged at token {k} ({b_tok} vs {t_tok}) with a " + f"{arm} top-2 logprob gap of {margin:.4f} > " + f"{_NEAR_TIE_LOGPROB_MARGIN}: not a near-tie flip" + ) + print( + f"prompt {prompt_idx}: accepted near-tie flip at token {k} " + f"({b_tok} vs {t_tok}); {k} tokens matched before it" + ) @skip_pre_hopper class TestKvPoolRebalanceAccuracy: - """Token-exact greedy-decode equivalence under rebalance. + """Greedy-decode equivalence under rebalance. Compares rebalance=off and rebalance=on with a forced mid-generation - adjust(). + adjust(). Outputs must match token for token, except that one divergence + at a near-tie is tolerated (see _assert_tokens_match_or_near_tie). """ MODEL_PATH = f"{llm_models_root()}/gemma/gemma-3-1b-it/" @@ -164,8 +242,4 @@ def test_rebalance_matches_baseline(self, disable_overlap, monkeypatch): assert len(baseline) == len(treated) == len(_PROMPTS) for i, (b, t) in enumerate(zip(baseline, treated)): - assert b == t, ( - f"prompt {i}: rebalance changed greedy-decode output\n" - f" baseline: {b[:16]}...\n" - f" treated: {t[:16]}..." - ) + _assert_tokens_match_or_near_tie(i, b, t)