Download code/tt_diffusion_planner/reference/pipeline.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 6.1 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/pipeline.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/reference/pipeline.py
-
curl -L -o pipeline.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/pipeline.py
6.1 kB
| # 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" | |
| 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 --------------------------------------------------------------------------------------------- | |
| 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)"}) | |
| def sample_path(name: str = "kashiwanoha_dense.npz") -> Path: | |
| return Path(__file__).resolve().parents[1] / "samples" / name | |