changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
8.58 kB
# SPDX-License-Identifier: Apache-2.0
"""Golden tensors of the CPU reference: per-module taps and final outputs of one planning scene.
Two products per scene (``code/scripts/ref_golden.py`` writes both):
- the **full goldens** (``<dir>/<scene>.npz``, 10-25 MB, kept OUT of the bundle, by default under
``research/diffusion-planner/goldens``): the raw inputs, the host features, every encoder tap (mixer taps compacted
to the valid entities, with their row indices), the encoding, for each of the 11 decoder evaluations its input
``x`` (all 321 agents, for teacher forcing), its time, its block taps and output (valid agents), the solver
iterates, ``final_x0``, the turn pool and logits, the post-processed outputs, and the port's rewrite constants
(adaLN tables, pad-relative island constants);
- the **small goldens** (``tests/goldens/<stem>.outputs.npz``, ~0.1 MB, shipped): ``final_x0`` of the valid agents,
the logits, the ego trajectory columns, the turn command and the ego row of every solver iterate.
Tap names follow ``reference.model.TAP_NAMES``; arrays of valid entities carry ``<name>.rows`` (indices into the full
tensor). Device tests compare TT replay outputs against these, never TT against TT.
"""
from __future__ import annotations
import datetime as _dt
from pathlib import Path
from typing import Any, Dict, Optional
import numpy as np
from . import config as C
from ..host import pipeline as hp
from ..host.postprocess import denoising_steps_ego
from ..host.solver import solver_plan
from ..ttaw.golden import TapRegistry, save_goldens
from . import rewrites as R
from .pipeline import ReferencePlanner
__all__ = ["scene_goldens", "small_goldens", "lite_goldens", "write_scene", "SMALL_KEYS"]
SMALL_KEYS = ("final_x0", "final_x0.rows", "logit", "trajectory", "turn_command", "denoising_ego")
MIXER_CATS = ("ego", "neighbor", "lane", "route", "polygon", "line_string")
def _rows(valid: np.ndarray) -> np.ndarray:
return np.flatnonzero(np.asarray(valid, bool)).astype(np.int32)
def scene_goldens(ref: ReferencePlanner, raw: Any, *, params: Optional[Dict[str, Any]] = None) -> Dict[str, np.ndarray]:
"""Run the reference on one scene and return the full golden dict (numpy)."""
taps = TapRegistry()
res = ref.run(raw, taps=taps, keep_eval_io=True)
prep = res.prepared
f = prep.features
g: Dict[str, np.ndarray] = {}
for k, v in prep.raw.items():
g[f"in.{k}"] = v
# host features (what the device consumes)
for name in ("ego", "neighbor", "neighbor_type", "static", "lane", "lane_attr", "lane_speed", "lane_has_speed",
"route", "route_attr", "route_speed", "route_has_speed", "polygon", "line_string", "goal",
"ego_shape", "turn", "token_valid", "key_valid", "pos"):
g[f"host.{name}"] = np.asarray(getattr(f, name))
for cat, v in f.valid.items():
g[f"host.valid.{cat}"] = np.asarray(v, bool)
g["host.agent_valid"] = prep.decoder.agent_valid
g["host.current_states"] = prep.decoder.current_states
t = taps.to_dict()
# encoder taps: mixer internals compacted to the valid entities
for cat in MIXER_CATS:
rows = _rows(f.valid[cat])
for suffix in ("pre", "mixer"):
g[f"enc.{cat}.{suffix}"] = t[f"enc.{cat}.{suffix}"][rows]
g[f"enc.{cat}.{suffix}.rows"] = rows
for name, _ in C.TOKEN_LAYOUT:
g[f"enc.{name}"] = t[f"enc.{name}"]
g["enc.tokens"] = t["enc.tokens"]
for i in range(C.FUSION_DEPTH):
g[f"enc.fusion.{i}"] = t[f"enc.fusion.{i}"]
g["enc.encoding"] = res.encoding
# decoder evaluations: inputs for teacher forcing, outputs and block taps on the valid agents
arows = _rows(prep.decoder.agent_valid)
g["dec.rows"] = arows
g["dec.t"] = np.asarray(res.eval_times, np.float32)
g["dec.x_in"] = np.stack(res.eval_inputs).astype(np.float32) # [11, 321, 81, 4]
g["dec.out"] = np.stack(res.eval_outputs)[:, arows].astype(np.float32)
for k in range(len(res.eval_times)):
g[f"dec.{k}.temb"] = t[f"dec.{k}.temb"][0]
g[f"dec.{k}.x"] = t[f"dec.{k}.x"][arows]
for i in range(C.DIT_DEPTH):
g[f"dec.{k}.block{i}"] = t[f"dec.{k}.block{i}"][arows]
g["solver.iterates"] = np.stack(res.denoising_steps).astype(np.float32)[:, arows]
g["solver.timesteps"] = np.asarray(res.denoising_timesteps, np.float32)
g["final_x0"] = res.final_x0
g["turn.pool"] = t["turn.pool"]
g["turn.logit"] = res.logit
# post-processed outputs (the API's Trajectory)
p = {k: spec[3] for k, spec in hp.RUNTIME_PARAMS.items()}
p.update(params or {})
out = hp.make_output(res.final_x0, res.logit, prep, ref.normalization, p, denoising_steps=res.denoising_steps)
g["out.trajectory"] = out.poses
g["out.predicted_agents"] = out.predicted_agents
g["out.turn_command"] = np.asarray(out.turn_indicator["command"], np.int32)
g["out.denoising_ego"] = denoising_steps_ego(np.stack(res.denoising_steps), *ref.normalization.state())
# the port's rewrite constants (built in float64, rounded once)
plan = solver_plan(C.DPM_SOLVER_STEPS)
tab = R.adaln_tables(ref.weights.params, plan.eval_times)
g["port.adaln.temb"] = tab.temb
for i, blk in enumerate(tab.blocks):
for k, v in blk.items():
g[f"port.adaln.block{i}.{k}"] = v
g["port.adaln.final_gamma"], g["port.adaln.final_beta"] = tab.final_gamma, tab.final_beta
for cat in ("ego", "neighbor"):
cst = R.island_constants(ref.weights.params, cat)
for k in ("gelu_b1", "c0", "t1_pad", "g_pad", "t2_pad", "w_t1_valid"):
g[f"port.island.{cat}.{k}"] = getattr(cst, k)
g["port.solver"] = np.asarray([[u.order, u.t, u.a, u.b, u.c, u.r0] for u in plan.updates], np.float64)
return g
LITE_PREFIXES = ("host.valid.", "host.token_valid", "host.agent_valid", "enc.", "dec.rows", "dec.t", "final_x0",
"turn.", "out.", "solver.timesteps")
def lite_goldens(g: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
"""The compact per-scene goldens of the public-data instants (~1 MB): validity, the encoder category outputs and
the encoding on valid rows, final x0 of the valid agents, logits and the post-processed outputs (no inputs: they
stay in ``public_data/inputs``; no mixer or per-evaluation decoder taps)."""
out: Dict[str, np.ndarray] = {}
tok = np.flatnonzero(g["host.token_valid"])
for k, v in g.items():
if not k.startswith(LITE_PREFIXES) or k.startswith(("enc.fusion.", "enc.tokens")):
continue
if k.endswith((".pre", ".mixer", ".pre.rows", ".mixer.rows")):
continue
out[k] = v
for name, _ in C.TOKEN_LAYOUT:
rows = np.flatnonzero(g[f"host.valid.{name}"])
out[f"enc.{name}"], out[f"enc.{name}.rows"] = g[f"enc.{name}"][rows], rows.astype(np.int32)
out["enc.encoding"], out["enc.encoding.rows"] = g["enc.encoding"][tok], tok.astype(np.int32)
out["final_x0"] = g["final_x0"][g["dec.rows"]]
return out
def small_goldens(g: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
"""The shipped subset of a full golden dict (see the module docstring)."""
rows = g["dec.rows"]
return {"final_x0": g["final_x0"][rows], "final_x0.rows": rows, "logit": g["turn.logit"],
"trajectory": g["out.trajectory"], "turn_command": g["out.turn_command"],
"denoising_ego": g["out.denoising_ego"]}
def write_scene(ref: ReferencePlanner, raw: Any, scene: str, full_dir: Optional[Path], small_dir: Optional[Path],
meta: Optional[Dict[str, Any]] = None, *, lite: bool = False) -> Dict[str, Any]:
g = scene_goldens(ref, raw)
info = {"scene": scene, "created": _dt.datetime.now(_dt.timezone.utc).isoformat(timespec="seconds"),
"weights_sha256": ref.weights.sha256, "solver_steps": C.DPM_SOLVER_STEPS,
"valid_counts": {k: int(np.asarray(v).sum()) for k, v in g.items() if k.startswith("host.valid.")},
"producer": "tt_diffusion_planner.reference (fp32 CPU, torch)", **(meta or {})}
paths = {}
if full_dir is not None:
tensors = lite_goldens(g) if lite else g
info["lite"] = bool(lite)
paths["full"] = str(save_goldens(Path(full_dir) / f"{scene}.npz", tensors, info, compress=True))
if small_dir is not None:
paths["small"] = str(save_goldens(Path(small_dir) / f"{scene}.outputs.npz", small_goldens(g), info,
compress=True))
return {"info": info, "paths": paths, "goldens": g}