File size: 2,147 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 | # SPDX-License-Identifier: Apache-2.0
"""CPU reference of Diffusion Planner v5.0 (Autoware diffusion_planner) -- the ground truth every PCC /
output-agreement gate of the TT port compares against. Importable without ttnn (torch, numpy, onnx only).
- ``config.py`` dimensions, token layout, constants and node parameters, each with its Autoware / ONNX source.
- ``weights.py`` the three ONNX files and ``diffusion_planner.param.json`` read as DATA with the vendored
``ttaw.weights.OnnxWeights``, every tensor addressed through its consuming node, into one canonical
``{name: float32}`` dict that the reference and the ttnn graph share (no BatchNorm to fold).
- ``model.py`` pure-PyTorch fp32 encoder / DiT decoder / turn head with per-module taps.
- ``rewrites.py`` the exact graph rewrites the TT port applies (per-step adaLN tables folded into the LayerNorm
affine, hoisted cross-attention K/V, the pad-relative fp32 pre-projection island) as float64-built
constants plus CPU forwards that use them, tested against ``model.py``.
- ``pipeline.py`` ``ReferencePlanner``: host pre-processing (``..host``) -> encoder -> DPM-Solver++(2M) loop over 11
decoder evaluations -> turn head -> host post-processing; ``run()`` records taps for goldens.
- ``ort.py`` ONNX Runtime on the shipped ONNX (the reference's oracle; research venv and tests only).
- ``goldens.py`` golden generation (taps + outputs per scene) and the small goldens kept in ``tests/goldens``.
The ``to_dict()`` of ``ReferencePlanner()(inputs=sample)`` is stored as ``samples/<stem>.reference.json``:
``server/smoke_test.py`` compares the served output with it.
"""
__all__ = ["ReferencePlanner", "load_weights", "find_weights_dir"]
def __getattr__(name):
if name == "ReferencePlanner":
from .pipeline import ReferencePlanner
return ReferencePlanner
if name in ("load_weights", "find_weights_dir"):
from . import weights
return getattr(weights, name)
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|