# SPDX-License-Identifier: Apache-2.0 """The fp32 CPU reference planner: host pre-processing -> encoder -> DPM-Solver++(2M) over the DiT decoder (11 evaluations) -> turn head -> host post-processing, i.e. the Autoware ``multi_step`` mode with guidance off. from tt_diffusion_planner.reference import ReferencePlanner ref = ReferencePlanner() # weights: find_weights_dir() or weights_dir=... out = ref(inputs="code/tt_diffusion_planner/samples/kashiwanoha_dense.npz") # ttaw.outputs.Trajectory res = ref.run(raw_inputs, taps=TapRegistry()) # every intermediate (golden generation) ``__call__`` returns exactly what the device model returns (same ``host.prepare`` / ``host.make_output``), so the stored ``samples/.reference.json`` is this class's ``to_dict()``. """ from __future__ import annotations import time from dataclasses import dataclass, field from pathlib import Path from typing import Any, Dict, List, Optional import numpy as np import torch from . import config as C from ..host import pipeline as hp from ..host.solver import apply_prefix_constraint, dpm_solver_sample from ..ttaw import io as tio from ..ttaw.golden import NULL_TAPS, TapRegistry from .model import Decoder, Encoder, TurnHead, torch_params from .weights import PlannerWeights, find_weights_dir, load_weights __all__ = ["ReferencePlanner", "ReferenceResult"] MODEL_NAME = "diffusion-planner-p150" @dataclass class ReferenceResult: prepared: hp.Prepared encoding: np.ndarray # [564, 256] final_x0: np.ndarray # [321, 81, 4] normalised logit: np.ndarray # [5] denoising_steps: List[np.ndarray] denoising_timesteps: List[float] eval_times: List[float] eval_inputs: List[np.ndarray] = field(default_factory=list) # x fed to each decoder evaluation eval_outputs: List[np.ndarray] = field(default_factory=list) # model_output of each evaluation timing_ms: Dict[str, float] = field(default_factory=dict) class ReferencePlanner: """CPU fp32 reference of the deployed network plus the shared host code. ``threads`` bounds torch's intra-op threads (the workspace rule is <= 4).""" def __init__(self, weights_dir: Optional[str] = None, *, weights: Optional[PlannerWeights] = None, threads: Optional[int] = 4, steps: int = C.DPM_SOLVER_STEPS): if weights is None: wd = find_weights_dir(weights_dir) if wd is None: raise FileNotFoundError("Diffusion Planner v5.0 weights not found: pass weights_dir= or set " "DIFFUSION_PLANNER_WEIGHTS_DIR") weights = load_weights(wd) self.weights = weights self.normalization = weights.normalization if threads: torch.set_num_threads(int(threads)) p = torch_params(weights.params) self.encoder, self.decoder, self.turn = Encoder(p), Decoder(p), TurnHead(p) self.steps = steps # ---- network --------------------------------------------------------------------------------------------- @torch.no_grad() def run(self, raw: Any, *, taps: TapRegistry = NULL_TAPS, keep_eval_io: bool = False) -> ReferenceResult: """One plan on raw ONNX-named inputs (mapping, ``.npz`` path / bytes or the JSON envelope).""" arrays = tio.load_named_arrays(raw, C.INPUT_SCHEMA) t0 = time.perf_counter() prep = hp.prepare(arrays, self.normalization.observation) t1 = time.perf_counter() encoding = self.encoder.forward(prep.features, taps) t2 = time.perf_counter() kv = self.decoder.cross_kv(encoding) cs = prep.decoder.current_states eval_in: List[np.ndarray] = [] eval_out: List[np.ndarray] = [] def model_fn(x: np.ndarray, t) -> np.ndarray: k = len(eval_out) y = self.decoder.forward(x, t, kv, prep.decoder.agent_valid, taps, prefix=f"dec.{k}").numpy() if keep_eval_io: eval_in.append(x.copy()) eval_out.append(y if keep_eval_io else np.empty(0)) return y res = dpm_solver_sample(prep.x_T, model_fn, lambda x: apply_prefix_constraint(x, cs), self.steps) t3 = time.perf_counter() logit = self.turn.forward(encoding, res.final_x, taps).numpy() t4 = time.perf_counter() for k, x in enumerate(res.denoising_steps): taps.tap(f"solver.x{k}", x) taps.tap("final_x0", res.final_x) return ReferenceResult(prep, encoding.numpy(), res.final_x, logit, res.denoising_steps, res.denoising_timesteps, res.eval_times, eval_in, eval_out if keep_eval_io else [], {"preprocess": (t1 - t0) * 1e3, "encoder": (t2 - t1) * 1e3, "solver": (t3 - t2) * 1e3, "turn": (t4 - t3) * 1e3}) # ---- the API-shaped call ---------------------------------------------------------------------------------- def __call__(self, inputs: Any, **params: Any): """``model(inputs=...)`` of the CPU reference: the same host pre / post-processing as the device model.""" p = {k: spec[3] for k, spec in hp.RUNTIME_PARAMS.items()} unknown = sorted(set(params) - set(p)) if unknown: raise tio.InputError(f"unknown parameter(s) {unknown}; allowed: {sorted(p)}") p.update(params) t0 = time.perf_counter() res = self.run(inputs) total = (time.perf_counter() - t0) * 1e3 steps = res.denoising_steps if p.get("return_denoising_steps") else None return hp.make_output(res.final_x0, res.logit, res.prepared, self.normalization, p, model=MODEL_NAME, denoising_steps=steps, timing_ms={**res.timing_ms, "total": total}, meta={"reference": "fp32 CPU (torch)"}) @staticmethod def sample_path(name: str = "kashiwanoha_dense.npz") -> Path: return Path(__file__).resolve().parents[1] / "samples" / name