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