changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
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
@dataclass
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)