File size: 7,129 Bytes
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
# 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}")