Skip to content
Closed
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
17 changes: 15 additions & 2 deletions families/qwen/tests/runtime_receipt.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,16 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Qwen native-KV runtime receipt checks used by family E2E."""
"""Qwen prompt and native-KV runtime receipt checks used by family E2E."""

from __future__ import annotations

import re

_PREFILL = re.compile(r"^\[trtmc\.prefill\] tokens=(\d+) launches=(\d+) max_chunk=(\d+)$")
_PREFILL = re.compile(
r"^(?:\[[^,\]]+,\d+\]<stderr>:)?\s*"
r"\[trtmc\.prefill\] tokens=(\d+) launches=(\d+) max_chunk=(\d+)$"
)
_RUNTIME_ERROR = re.compile(
r"\[trt\]\s+ERROR:|IExecutionContext::enqueueV3:\s+Error Code|"
r"Internal Error:|Cuda Runtime|illegal memory access",
Expand All @@ -24,6 +27,16 @@ def prefill_observations(stderr: str) -> tuple[tuple[int, int, int], ...]:
return tuple(values)


def assert_prompt_token_count(payload: dict, prompt_tokens: int) -> None:
observations = prefill_observations(str(payload["runtime_stderr"]))
assert observations, "native execution did not report its prompt token count"
# Tensor-parallel ranks each report the complete prompt, not disjoint pieces.
for tokens, _, _ in observations:
assert tokens == prompt_tokens, (
f"native prompt has {tokens} tokens; reference has {prompt_tokens}"
)


def assert_native_kv_receipt(payload: dict, case: dict, prompt_tokens: int) -> None:
stderr = str(payload["runtime_stderr"])
expected_rows = int(case["expected_kv_cache_rows"])
Expand Down
11 changes: 8 additions & 3 deletions families/qwen/tests/test_e2e.py
Original file line number Diff line number Diff line change
Expand Up @@ -278,15 +278,15 @@ def _run_native(
return payload


def _raw_prompt_token_count(model_dir: Path, manifest: dict, prompt: str) -> int:
def _prompt_token_count(model_dir: Path, manifest: dict, case: dict, prompt: str) -> int:
from transformers import AutoTokenizer

tokenizer = AutoTokenizer.from_pretrained(
model_dir,
local_files_only=True,
trust_remote_code=bool(manifest.get("trust_remote_code", False)),
)
return len(tokenizer.encode(prompt, add_special_tokens=False))
return int(_render_prompt(tokenizer, prompt, case)["input_ids"].shape[-1])


@cache
Expand Down Expand Up @@ -765,7 +765,10 @@ def test_e2e(case_name: str, request, tmp_path: Path) -> None:
)
if manifest["task"] == "embedding":
return _embedding_e2e(manifest, case, model_dir, runtime_root, torch, tmp_path)
from families.qwen.tests.runtime_receipt import assert_prompt_token_count

prompt = _prompt(case)
prompt_tokens = _prompt_token_count(model_dir, manifest, case, prompt)
record_evidence("inputs", {"prompt": prompt})
bundle = tmp_path / manifest["bundle"]

Expand All @@ -784,6 +787,8 @@ def test_e2e(case_name: str, request, tmp_path: Path) -> None:
tmp_path,
)
record_evidence("native", payload)
with evidence_stage("compare"):
assert_prompt_token_count(payload, prompt_tokens)
if case_name in _LOGIT_ORACLES:
with evidence_stage("native"):
payload["logits_trace"] = _native_logits_trace(bundle, prompt, case, tp_size, tmp_path)
Expand All @@ -802,6 +807,7 @@ def test_e2e(case_name: str, request, tmp_path: Path) -> None:
)
record_evidence("native", repeated)
with evidence_stage("compare"):
assert_prompt_token_count(repeated, prompt_tokens)
assert repeated["token_ids"] == payload["token_ids"]
with evidence_stage("compare"):
assert repeated["text"] == payload["text"]
Expand All @@ -817,7 +823,6 @@ def test_e2e(case_name: str, request, tmp_path: Path) -> None:
from families.qwen.tests.runtime_receipt import assert_native_kv_receipt

assert "expected_prompt_tokens" in case
prompt_tokens = _raw_prompt_token_count(model_dir, manifest, prompt)
with evidence_stage("compare"):
assert_native_kv_receipt(payload, case, prompt_tokens)
record_evidence("reference", {"mode": "contract_only", "oracle": "family runtime contract"})
Expand Down
33 changes: 32 additions & 1 deletion families/qwen/tests/test_runtime_receipt.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@

import pytest

from families.qwen.tests.runtime_receipt import assert_native_kv_receipt
from families.qwen.tests.runtime_receipt import assert_native_kv_receipt, assert_prompt_token_count


def _case() -> dict:
Expand Down Expand Up @@ -60,3 +60,34 @@ def test_long_regression_keeps_exact_prompt_and_observed_token_gates() -> None:
assert_native_kv_receipt(_payload(stderr), case, 65)
with pytest.raises(AssertionError):
assert_native_kv_receipt(_payload(stderr), case, 64)


@pytest.mark.parametrize("tokens", [20, 24, 144])
def test_prompt_count_matches_without_kv_expectations(tokens: int) -> None:
assert_prompt_token_count(
_payload(f"[trtmc.prefill] tokens={tokens} launches=1 max_chunk={tokens}"), tokens
)


def test_prompt_count_rejects_extra_thinking_prefix() -> None:
with pytest.raises(AssertionError, match="native prompt has 24 tokens; reference has 20"):
assert_prompt_token_count(
_payload("[trtmc.prefill] tokens=24 launches=1 max_chunk=24"), 20
)


def test_prompt_count_requires_native_receipt() -> None:
with pytest.raises(AssertionError, match="did not report"):
assert_prompt_token_count(_payload(""), 20)


@pytest.mark.parametrize("tagged", [False, True])
def test_prompt_count_checks_each_tensor_parallel_rank(tagged: bool) -> None:
receipt = "[trtmc.prefill] tokens=20 launches=1 max_chunk=20"
rank0 = "[1,0]<stderr>:" if tagged else ""
rank1 = "[1,1]<stderr>:" if tagged else ""
assert_prompt_token_count(_payload(f"{rank0}{receipt}\n{rank1}{receipt}"), 20)
with pytest.raises(AssertionError, match="native prompt has 24 tokens"):
assert_prompt_token_count(
_payload(f"{rank0}{receipt}\n{rank1}[trtmc.prefill] tokens=24 launches=1 max_chunk=24"), 20
)
Loading