Download code/tt_diffusion_planner/reference/goldens.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.58 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/goldens.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/reference/goldens.py
-
curl -L -o goldens.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/goldens.py
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} | |