Download code/tt_diffusion_planner/reference/ort.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 6.72 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/ort.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/reference/ort.py
-
curl -L -o ort.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/reference/ort.py
6.72 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """ONNX Runtime on the shipped v5.0 ONNX files: the oracle of the CPU reference (research and tests only; | |
| onnxruntime is not a runtime dependency of the bundle or the image). | |
| ``OrtPlanner.run(raw)`` is the Autoware ``multi_step`` pipeline with ORT as the backend (a port of the verified | |
| research script ``research/diffusion-planner/scripts/dp_reference.py``): the node's normalization and speed masks, | |
| encoder.onnx, the DPM-Solver++(2M) loop of :mod:`tt_diffusion_planner.host.solver` calling decoder.onnx 11 times, | |
| then turn_indicator.onnx. ``taps=[...]`` exposes intermediate tensors of the encoder / decoder graphs by ONNX tensor | |
| name (the decoder taps are recorded per evaluation), for the per-module PCC test of the reference. | |
| """ | |
| from __future__ import annotations | |
| from dataclasses import dataclass, field | |
| from pathlib import Path | |
| from typing import Any, Dict, List, Sequence | |
| import numpy as np | |
| from . import config as C | |
| from ..host.normalize import normalize_inputs, speed_masks | |
| from ..host.solver import apply_prefix_constraint, dpm_solver_sample | |
| from ..host.features import decoder_masks | |
| from ..ttaw import io as tio | |
| from .weights import load_param_json | |
| __all__ = ["OrtPlanner", "OrtResult", "available"] | |
| def available() -> bool: | |
| import importlib.util | |
| return importlib.util.find_spec("onnxruntime") is not None | |
| def _session(model: Any, threads: int, optimize: bool = True): | |
| import onnxruntime as ort | |
| so = ort.SessionOptions() | |
| so.graph_optimization_level = (ort.GraphOptimizationLevel.ORT_ENABLE_ALL if optimize | |
| else ort.GraphOptimizationLevel.ORT_DISABLE_ALL) | |
| so.intra_op_num_threads = threads | |
| so.inter_op_num_threads = 1 | |
| so.log_severity_level = 3 | |
| src = model.SerializeToString() if hasattr(model, "SerializeToString") else str(model) | |
| return ort.InferenceSession(src, so, providers=["CPUExecutionProvider"]) | |
| def _with_outputs(path: Path, names: Sequence[str]): | |
| """The model with extra graph outputs (intermediate tensors by name).""" | |
| import onnx | |
| m = onnx.load(str(path)) | |
| have = {o.name for o in m.graph.output} | |
| m.graph.output.extend([onnx.ValueInfoProto(name=n) for n in names if n not in have]) | |
| return m | |
| class OrtResult: | |
| norm: Dict[str, np.ndarray] | |
| encoding: np.ndarray # [1, 564, 256] | |
| final_x0: np.ndarray # [1, 321, 81, 4] | |
| logit: np.ndarray # [1, 5] | |
| denoising_steps: List[np.ndarray] | |
| denoising_timesteps: List[float] | |
| eval_times: List[float] | |
| eval_inputs: List[np.ndarray] = field(default_factory=list) | |
| eval_outputs: List[np.ndarray] = field(default_factory=list) | |
| taps: Dict[str, np.ndarray] = field(default_factory=dict) | |
| class OrtPlanner: | |
| """``OrtPlanner(weights_dir, threads=4, encoder_taps=(...), decoder_taps=(...))``. With taps the sessions run | |
| without graph optimizations (every intermediate kept as computed); without, ``ORT_ENABLE_ALL`` like | |
| ``onnxruntime_inference.cpp:172``.""" | |
| def __init__(self, weights_dir: Path, *, threads: int = 4, encoder_taps: Sequence[str] = (), | |
| decoder_taps: Sequence[str] = (), turn_taps: Sequence[str] = ()): | |
| wd = Path(weights_dir) | |
| self.normalization = load_param_json(wd / C.PARAM_JSON) | |
| self.encoder_taps, self.decoder_taps, self.turn_taps = list(encoder_taps), list(decoder_taps), list(turn_taps) | |
| opt = not (encoder_taps or decoder_taps or turn_taps) | |
| self.enc = _session(_with_outputs(wd / C.ENCODER_ONNX, self.encoder_taps), threads, opt) | |
| self.dec = _session(_with_outputs(wd / C.DECODER_ONNX, self.decoder_taps), threads, opt) | |
| self.turn = _session(_with_outputs(wd / C.TURN_INDICATOR_ONNX, self.turn_taps), threads, opt) | |
| def normalize(self, raw: Any) -> Dict[str, np.ndarray]: | |
| arrays = tio.load_named_arrays(raw, C.INPUT_SCHEMA) | |
| norm = normalize_inputs(arrays, self.normalization.observation) | |
| norm.update(speed_masks(norm)) | |
| return norm | |
| def encode(self, norm: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]: | |
| feed = {k: norm[k] for k in C.ENCODER_INPUTS} | |
| outs = self.enc.run(["encoding"] + self.encoder_taps, feed) | |
| return dict(zip(["encoding"] + self.encoder_taps, outs)) | |
| def decode(self, encoding: np.ndarray, x: np.ndarray, t, neighbor_agents_past: np.ndarray) -> Dict[str, Any]: | |
| dt = np.full((1, C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, 1), np.float32(t), np.float32) | |
| feed = {"encoding": encoding, "sampled_trajectories": np.asarray(x, np.float32).reshape( | |
| 1, C.MAX_NUM_AGENTS, C.OUTPUT_T + 1, C.POSE_DIM), "diffusion_time": dt, | |
| "neighbor_agents_past": neighbor_agents_past} | |
| outs = self.dec.run(["model_output"] + self.decoder_taps, feed) | |
| return dict(zip(["model_output"] + self.decoder_taps, outs)) | |
| def turn_logit(self, encoding: np.ndarray, final_x0: np.ndarray) -> Dict[str, np.ndarray]: | |
| outs = self.turn.run(["turn_indicator_logit"] + self.turn_taps, | |
| {"encoding": encoding, "final_x0": np.asarray(final_x0, np.float32)}) | |
| return dict(zip(["turn_indicator_logit"] + self.turn_taps, outs)) | |
| def run(self, raw: Any, *, steps: int = C.DPM_SOLVER_STEPS, keep_eval_io: bool = False) -> OrtResult: | |
| norm = self.normalize(raw) | |
| enc = self.encode(norm) | |
| encoding = enc["encoding"] | |
| taps: Dict[str, np.ndarray] = {f"enc:{k}": v for k, v in enc.items() if k != "encoding"} | |
| cs = decoder_masks(norm).current_states | |
| nb = norm["neighbor_agents_past"] | |
| eval_in: List[np.ndarray] = [] | |
| eval_out: List[np.ndarray] = [] | |
| def model_fn(x: np.ndarray, t) -> np.ndarray: | |
| k = len(eval_out) | |
| out = self.decode(encoding, x, t, nb) | |
| for name in self.decoder_taps: | |
| taps[f"dec{k}:{name}"] = out[name] | |
| y = out["model_output"][0] | |
| eval_out.append(y if keep_eval_io else np.empty(0)) | |
| if keep_eval_io: | |
| eval_in.append(np.array(x, np.float32, copy=True)) | |
| return y | |
| res = dpm_solver_sample(norm["sampled_trajectories"][0], model_fn, | |
| lambda x: apply_prefix_constraint(x, cs), steps) | |
| tl = self.turn_logit(encoding, res.final_x[None]) | |
| taps.update({f"turn:{k}": v for k, v in tl.items() if k != "turn_indicator_logit"}) | |
| return OrtResult(norm, encoding, res.final_x[None], tl["turn_indicator_logit"], res.denoising_steps, | |
| res.denoising_timesteps, res.eval_times, eval_in, eval_out if keep_eval_io else [], taps) | |