Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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/"
Expand All @@ -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)
Loading