changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
7.13 kB
# SPDX-License-Identifier: Apache-2.0
"""Per-module PCC gates: each device stage vs the fp32 CPU reference on the same inputs (needs one p150).
bin/devrun -t 1800 -- python -m pytest -q -s code/tt_diffusion_planner/tests/test_pcc_device.py
Gates (PLAN.md 2.12, SPEC 6.3): every encoder module on its VALID rows (entities / tokens) PCC >= 0.999, and every
single decoder evaluation, teacher-forced on the reference's own solver inputs and encoding, PCC >= 0.999 (the gate
value is the minimum over the 11 evaluations and the scenes; the t = 0 slot is excluded: the port zeroes it in the
last projection, an exact rewrite because the prefix constraint overwrites it). They are checked on trace REPLAY
outputs of the debug variants of ``tt.model.TtDiffusionPlanner`` (``encoder_taps`` / ``decode_once``), never TT
against TT; ``test_plan_replay_equals_eager`` checks the production ``plan`` trace bit for bit against an eager run.
``GATES`` (vendored ``ttaw.golden.GateRegistry``) freezes each gate into ``test_pcc_device.gates.json`` next to this
file at its first green run and refuses a looser declaration afterwards (PLAN.md 0.3; ``TTAW_GATES_READONLY=1`` for
verification runs). Changing a frozen gate means editing that JSON by hand and disclosing it in VERIFICATION_<date>.md.
Goldens: ``research/diffusion-planner/goldens/<scene>.npz`` when present (``code/scripts/ref_golden.py``), else the CPU
reference computes them on the fly (``reference.goldens.scene_goldens``).
"""
from __future__ import annotations
from pathlib import Path
import numpy as np
import pytest
from tt_diffusion_planner.reference import config as C
from tt_diffusion_planner.ttaw.golden import GateRegistry
from tt_diffusion_planner.ttaw.metrics import pcc
pytestmark = pytest.mark.device
PKG = Path(__file__).resolve().parents[1]
FULL_GOLDENS = PKG.parents[3] / "research" / "diffusion-planner" / "goldens"
SAMPLES = ("kashiwanoha_dense", "straight_road") # every module has valid rows in at least one of them
# stage -> minimum PCC vs the fp32 reference on valid rows (keep the names stable across rounds)
GATES = GateRegistry.for_test(__file__, {
"enc.ego": 0.999, "enc.neighbor": 0.999, "enc.lane": 0.999, "enc.route": 0.999, "enc.polygon": 0.999,
"enc.line_string": 0.999, "enc.goal": 0.999, "enc.ego_shape": 0.999, "enc.turn": 0.999,
"enc.encoding": 0.999, "dec.eval": 0.999})
# reported, not gated (diagnostics for the bring-up of each module)
DIAGNOSTICS = (tuple(f"enc.{c}.{p}" for c in ("ego", "neighbor", "lane", "route", "polygon", "line_string")
for p in ("pre", "mixer")) + ("enc.tokens",) + tuple(f"enc.fusion.{i}" for i in range(6)))
@pytest.fixture(scope="module")
def weights():
from tt_diffusion_planner.reference.weights import find_weights_dir, load_weights
wd = find_weights_dir()
if wd is None:
pytest.skip("weights not found")
return load_weights(wd)
@pytest.fixture(scope="module")
def planner(device, weights):
from tt_diffusion_planner.tt.model import TtDiffusionPlanner
tt = TtDiffusionPlanner(device, weights, debug=True)
tt.capture()
yield tt
tt.release()
def _goldens(stem: str):
path = FULL_GOLDENS / f"{stem}.npz"
if path.is_file():
with np.load(path, allow_pickle=False) as z:
return {k: z[k] for k in z.files if k != "__meta__"}
from tt_diffusion_planner.reference.goldens import scene_goldens
from tt_diffusion_planner.reference.pipeline import ReferencePlanner
return scene_goldens(ReferencePlanner(threads=4), str(PKG / "samples" / f"{stem}.npz"))
@pytest.fixture(scope="module")
def scenes(weights):
from tt_diffusion_planner.host import pipeline as hp
out = []
for stem in SAMPLES:
g = _goldens(stem)
raw = {k: g[f"in.{k}"] for k in C.INPUT_NAMES}
out.append((stem, hp.prepare(raw, weights.normalization.observation), g))
return out
@pytest.fixture(scope="module")
def stage_outputs(planner, scenes):
"""``{stage: [(device, reference), ...]}`` over the samples, valid rows only (trace replays)."""
out: dict = {}
for stem, prep, g in scenes:
taps = planner.encoder_taps(prep) # trace replay of the encoder debug variant
for name, _ in C.TOKEN_LAYOUT:
rows = np.flatnonzero(g[f"host.valid.{name}"])
if rows.size:
out.setdefault(f"enc.{name}", []).append((taps[f"enc.{name}"][rows], g[f"enc.{name}"][rows]))
tok = np.flatnonzero(g["host.token_valid"])
out.setdefault("enc.encoding", []).append((taps["enc.encoding"][tok], g["enc.encoding"][tok]))
out.setdefault("enc.tokens", []).append((taps["enc.tokens"][tok], g["enc.tokens"][tok]))
for i in range(C.FUSION_DEPTH):
out.setdefault(f"enc.fusion.{i}", []).append((taps[f"enc.fusion.{i}"][tok], g[f"enc.fusion.{i}"][tok]))
for c in ("ego", "neighbor", "lane", "route", "polygon", "line_string"):
rows = g[f"enc.{c}.pre.rows"]
for part in ("pre", "mixer"):
if rows.size:
out.setdefault(f"enc.{c}.{part}", []).append((taps[f"enc.{c}.{part}"][rows], g[f"enc.{c}.{part}"]))
arows = g["dec.rows"]
for k, t in enumerate(g["dec.t"]):
got = planner.decode_once(prep, g["dec.x_in"][k], float(t), encoding=g["enc.encoding"])
out.setdefault("dec.eval", []).append((got[arows][:, 1:], g["dec.out"][k][:, 1:]))
out.setdefault(f"dec.eval.{k}", []).append((got[arows][:, 1:], g["dec.out"][k][:, 1:]))
return out
@pytest.mark.parametrize("stage", GATES.names())
def test_stage_pcc(stage_outputs, stage):
pairs = stage_outputs.get(stage)
assert pairs, f"no valid rows for {stage} in {SAMPLES}"
value = min(pcc(dev, ref) for dev, ref in pairs)
print(GATES.require(stage, value))
@pytest.mark.parametrize("stage", DIAGNOSTICS + tuple(f"dec.eval.{k}" for k in range(C.DPM_SOLVER_STEPS + 1)))
def test_stage_diagnostics(stage_outputs, stage):
for dev, ref in stage_outputs.get(stage, []):
print(stage, f"pcc={pcc(dev, ref):.6f}", f"max_abs={float(np.abs(np.asarray(dev) - ref).max()):.3e}")
def test_plan_replay_equals_eager(planner, scenes):
"""The production ``plan`` trace: replay == an eager run of the same graph on the same inputs, bit for bit, and
the prefix constraint holds exactly."""
for stem, prep, g in scenes:
replay = planner.forward(prep)
eager = planner.forward(prep, eager=True)
for key in ("final_x0", "logit"):
np.testing.assert_array_equal(replay[key], eager[key], err_msg=f"{stem}: {key}")
for a, b in zip(replay["denoising_steps"], eager["denoising_steps"]):
np.testing.assert_array_equal(a, b)
np.testing.assert_array_equal(replay["final_x0"][:, 0], prep.decoder.current_states)
rows = g["dec.rows"]
print(stem, f"final_x0 valid-agent pcc={pcc(replay['final_x0'][rows], g['final_x0'][rows]):.6f}",
f"logit max abs={float(np.abs(replay['logit'] - g['turn.logit']).max()):.4f}")