Skip to content
Open
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
306 changes: 163 additions & 143 deletions tests/v1/determinism/test_batch_invariance.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
TEST_MODEL,
_extract_step_logprobs,
_random_prompt,
get_attention_config,
skip_if_not_cuda,
skip_unsupported,
)
Expand Down Expand Up @@ -61,7 +62,7 @@ def test_v1_generation_is_deterministic_across_batch_sizes_with_needle(
seed = int(os.getenv("VLLM_TEST_SEED", "12345"))
random.seed(seed)

attention_config = {"backend": backend}
attention_config = get_attention_config(backend)
# Force the C++ RMSNorm implementation so we actually exercise the
# num_tokens-dependent block-size branches.
kernel_config = None
Expand Down Expand Up @@ -197,159 +198,128 @@ def test_logprobs_bitwise_batch_invariance_bs1_vs_bsN(
print(f"BATCH INVARIANCE MODE: Disabling custom all-reduce (TP={tp_size})")
print(f"{'=' * 80}\n")

llm = LLM(
model=TEST_MODEL,
tensor_parallel_size=tp_size,
max_num_seqs=128,
max_model_len=8192,
dtype="auto", # not everything is supported
gpu_memory_utilization=0.9,
attention_config={
"backend": backend,
"flex_attn_block_m": block_m,
"flex_attn_block_n": block_n,
},
)

# Use more realistic prompts for better token generation
prompts = [_random_prompt(10, 50) for _ in range(32)]

# TODO: Update prompts to have ragged lengths in order to test chunked prefill
# The above tests are not currently long enough to exercise chunking.
# prompts = (
# [_random_prompt(10, 50) for _ in range(28)]
# + [_random_prompt(256, 512) for _ in range(50)]
# + [_random_prompt(2048, 4096) for _ in range(50)]
# )

sp = SamplingParams(
temperature=0.6,
top_p=1.0,
max_tokens=16,
seed=1234,
logprobs=5,
)

# BS=1: run prompts individually and collect logprobs per step.
print("\n" + "=" * 80)
print("STARTING BS=1 RUNS (each prompt individually)")
print("=" * 80 + "\n")

bs1_logprobs_per_prompt = []
bs1_tokens_per_prompt = []
for idx, p in enumerate(prompts):
print(f"\n[BS=1] Running prompt {idx}/{len(prompts)} - Preview: {p[:80]}...")
outs = llm.generate([p], sp, use_tqdm=False)
assert len(outs) == 1
step_logprobs, token_ids = _extract_step_logprobs(outs[0])
if step_logprobs is None:
pytest.skip(
"Logits are not available on RequestOutput; "
"enable logprobs return to run this test."
)
bs1_logprobs_per_prompt.append(step_logprobs)
bs1_tokens_per_prompt.append(token_ids)
print(f"[BS=1] Prompt {idx} generated tokens: {token_ids}")

# BS=N: run prompts in a batch and collect logprobs per step for each
# prompt.
print("\n" + "=" * 80)
print(f"STARTING BS={len(prompts)} RUN (all prompts batched)")
print("=" * 80 + "\n")
_attn_cfg = {
**get_attention_config(backend),
**(
{"flex_attn_block_m": block_m, "flex_attn_block_n": block_n}
if backend != "GDN_ATTN"
else {}
),
}
llm = None
prompts: list[str] = []
failed_prompts: list[dict] = []
try:
llm = LLM(
model=TEST_MODEL,
tensor_parallel_size=tp_size,
max_num_seqs=128,
max_model_len=8192,
dtype="auto", # not everything is supported
gpu_memory_utilization=0.9,
enforce_eager=backend == "GDN_ATTN",
attention_config=_attn_cfg,
)

outs_batched = llm.generate(prompts, sp, use_tqdm=False)
assert len(outs_batched) == len(prompts)
bsN_logprobs_per_prompt = []
bsN_tokens_per_prompt = []
# Use more realistic prompts for better token generation
prompts = [_random_prompt(10, 50) for _ in range(32)]

# TODO: Update prompts to have ragged lengths in order to test chunked prefill
# The above tests are not currently long enough to exercise chunking.
# prompts = (
# [_random_prompt(10, 50) for _ in range(28)]
# + [_random_prompt(256, 512) for _ in range(50)]
# + [_random_prompt(2048, 4096) for _ in range(50)]
# )

sp = SamplingParams(
temperature=0.6,
top_p=1.0,
max_tokens=16,
seed=1234,
logprobs=5,
)

print(f"\n[BS={len(prompts)}] Processing batched outputs...")
for idx, o in enumerate(outs_batched):
tokens = o.outputs[0].token_ids if o.outputs else "N/A"
print(f"[BS={len(prompts)}] Prompt {idx} generated tokens: {tokens}")
step_logprobs, token_ids = _extract_step_logprobs(o)
if step_logprobs is None:
pytest.skip(
"Logits are not available on RequestOutput; "
"enable logprobs return to run this test."
)
bsN_logprobs_per_prompt.append(step_logprobs)
bsN_tokens_per_prompt.append(token_ids)
# BS=1: run prompts individually and collect logprobs per step.
print("\n" + "=" * 80)
print("STARTING BS=1 RUNS (each prompt individually)")
print("=" * 80 + "\n")

# Compare step-by-step logprobs for each prompt between BS=1 and BS=N runs.
failed_prompts = []
for i, (logprobs_bs1, logprobs_bsN, tokens_bs1, tokens_bsN) in enumerate(
zip(
bs1_logprobs_per_prompt,
bsN_logprobs_per_prompt,
bs1_tokens_per_prompt,
bsN_tokens_per_prompt,
)
):
if len(logprobs_bs1) != len(logprobs_bsN):
reason = (
f"Different number of steps: {len(logprobs_bs1)} (BS=1) "
f"vs {len(logprobs_bsN)} (BS=N)"
)
failed_prompts.append(
{
"prompt_idx": i,
"step": "all",
"reason": reason,
"prompt_preview": prompts[i][:100],
"bs1_tokens": tokens_bs1,
"bsN_tokens": tokens_bsN,
}
bs1_logprobs_per_prompt = []
bs1_tokens_per_prompt = []
for idx, p in enumerate(prompts):
print(
f"\n[BS=1] Running prompt {idx}/{len(prompts)} - Preview: {p[:80]}..."
)
continue

# Check if tokens match first
if tokens_bs1 != tokens_bsN:
failed_prompts.append(
{
"prompt_idx": i,
"step": "sampling",
"reason": "Different tokens sampled",
"prompt_preview": prompts[i][:100],
"bs1_tokens": tokens_bs1,
"bsN_tokens": tokens_bsN,
"bs1_all_logprobs": [
logprobs_bs1[s].tolist() for s in range(len(logprobs_bs1))
],
"bsN_all_logprobs": [
logprobs_bsN[s].tolist() for s in range(len(logprobs_bsN))
],
}
outs = llm.generate([p], sp, use_tqdm=False)
assert len(outs) == 1
step_logprobs, token_ids = _extract_step_logprobs(outs[0])
if step_logprobs is None:
pytest.skip(
"Logits are not available on RequestOutput; "
"enable logprobs return to run this test."
)
bs1_logprobs_per_prompt.append(step_logprobs)
bs1_tokens_per_prompt.append(token_ids)
print(f"[BS=1] Prompt {idx} generated tokens: {token_ids}")

# BS=N: run prompts in a batch and collect logprobs per step for each
# prompt.
print("\n" + "=" * 80)
print(f"STARTING BS={len(prompts)} RUN (all prompts batched)")
print("=" * 80 + "\n")

outs_batched = llm.generate(prompts, sp, use_tqdm=False)
assert len(outs_batched) == len(prompts)
bsN_logprobs_per_prompt = []
bsN_tokens_per_prompt = []

print(f"\n[BS={len(prompts)}] Processing batched outputs...")
for idx, o in enumerate(outs_batched):
tokens = o.outputs[0].token_ids if o.outputs else "N/A"
print(f"[BS={len(prompts)}] Prompt {idx} generated tokens: {tokens}")
step_logprobs, token_ids = _extract_step_logprobs(o)
if step_logprobs is None:
pytest.skip(
"Logits are not available on RequestOutput; "
"enable logprobs return to run this test."
)
bsN_logprobs_per_prompt.append(step_logprobs)
bsN_tokens_per_prompt.append(token_ids)

# Compare step-by-step logprobs for each prompt between BS=1 and BS=N runs.
for i, (logprobs_bs1, logprobs_bsN, tokens_bs1, tokens_bsN) in enumerate(
zip(
bs1_logprobs_per_prompt,
bsN_logprobs_per_prompt,
bs1_tokens_per_prompt,
bsN_tokens_per_prompt,
)
continue

for t, (a, b) in enumerate(zip(logprobs_bs1, logprobs_bsN)):
if a.shape != b.shape:
):
if len(logprobs_bs1) != len(logprobs_bsN):
reason = (
f"Different number of steps: {len(logprobs_bs1)} (BS=1) "
f"vs {len(logprobs_bsN)} (BS=N)"
)
failed_prompts.append(
{
"prompt_idx": i,
"step": t,
"reason": f"Shape mismatch: {a.shape} vs {b.shape}",
"step": "all",
"reason": reason,
"prompt_preview": prompts[i][:100],
"bs1_tokens": tokens_bs1,
"bsN_tokens": tokens_bsN,
}
)
break
continue

if not torch.equal(a, b):
max_diff = torch.abs(a - b).max().item()
# Print which token failed
print(f"\n[DIVERGENCE] Prompt {i}, Token {t}: max_diff={max_diff:.6e}")
bs1_tok = tokens_bs1[t] if t < len(tokens_bs1) else "N/A"
bsN_tok = tokens_bsN[t] if t < len(tokens_bsN) else "N/A"
print(f" Token IDs: bs1={bs1_tok}, bsN={bsN_tok}")
print(f" BS=1 logprob: {a.tolist()}")
print(f" BS=N logprob: {b.tolist()}")
# Check if tokens match first
if tokens_bs1 != tokens_bsN:
failed_prompts.append(
{
"prompt_idx": i,
"step": t,
"reason": f"Bitwise mismatch (max_diff={max_diff:.6e})",
"step": "sampling",
"reason": "Different tokens sampled",
"prompt_preview": prompts[i][:100],
"bs1_tokens": tokens_bs1,
"bsN_tokens": tokens_bsN,
Expand All @@ -361,9 +331,58 @@ def test_logprobs_bitwise_batch_invariance_bs1_vs_bsN(
],
}
)
break
continue

for t, (a, b) in enumerate(zip(logprobs_bs1, logprobs_bsN)):
if a.shape != b.shape:
failed_prompts.append(
{
"prompt_idx": i,
"step": t,
"reason": f"Shape mismatch: {a.shape} vs {b.shape}",
"prompt_preview": prompts[i][:100],
"bs1_tokens": tokens_bs1,
"bsN_tokens": tokens_bsN,
}
)
break

if not torch.equal(a, b):
max_diff = torch.abs(a - b).max().item()
# Print which token failed
print(
f"\n[DIVERGENCE] Prompt {i}, Token {t}: max_diff={max_diff:.6e}"
)
bs1_tok = tokens_bs1[t] if t < len(tokens_bs1) else "N/A"
bsN_tok = tokens_bsN[t] if t < len(tokens_bsN) else "N/A"
print(f" Token IDs: bs1={bs1_tok}, bsN={bsN_tok}")
print(f" BS=1 logprob: {a.tolist()}")
print(f" BS=N logprob: {b.tolist()}")
failed_prompts.append(
{
"prompt_idx": i,
"step": t,
"reason": f"Bitwise mismatch (max_diff={max_diff:.6e})",
"prompt_preview": prompts[i][:100],
"bs1_tokens": tokens_bs1,
"bsN_tokens": tokens_bsN,
"bs1_all_logprobs": [
logprobs_bs1[s].tolist()
for s in range(len(logprobs_bs1))
],
"bsN_all_logprobs": [
logprobs_bsN[s].tolist()
for s in range(len(logprobs_bsN))
],
}
)
break
finally:
with contextlib.suppress(Exception):
if llm is not None:
llm.shutdown()

# Print summary of all failures
# Print summary of all failures (after LLM shutdown so GPU memory is freed).
if failed_prompts:
print(f"\n{'=' * 80}")
fail_msg = (
Expand Down Expand Up @@ -392,7 +411,6 @@ def test_logprobs_bitwise_batch_invariance_bs1_vs_bsN(
print(f" Step {step_idx}: {logprobs}")
print(f"{'=' * 80}\n")

# Fail the test with summary
msg = (
f"Batch invariance violated in {len(failed_prompts)}/"
f"{len(prompts)} prompts. See output above for details."
Expand Down Expand Up @@ -420,7 +438,8 @@ def test_simple_generation(backend):
max_model_len=2048,
dtype="auto",
enable_prefix_caching=False,
attention_config={"backend": backend},
enforce_eager=backend == "GDN_ATTN",
attention_config=get_attention_config(backend),
)

prompt = "the capital of france is"
Expand Down Expand Up @@ -484,7 +503,8 @@ def test_logprobs_without_batch_invariance_should_fail(
max_num_seqs=32,
max_model_len=8192,
dtype="auto",
attention_config={"backend": backend},
enforce_eager=backend == "GDN_ATTN",
attention_config=get_attention_config(backend),
)

# build ragged prompts to change shapes significantly across BS=1 vs BS=N
Expand Down Expand Up @@ -703,7 +723,7 @@ def test_decode_logprobs_match_prefill_logprobs(
max_num_seqs=32,
max_model_len=8192,
dtype="auto",
attention_config={"backend": backend},
attention_config=get_attention_config(backend),
)

# Use a few test prompts
Expand Down
Loading