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}")