File size: 6,097 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 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 | # 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/<stem>.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
|