Download code/tt_diffusion_planner/host/solver.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 11.3 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/host/solver.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/host/solver.py
-
curl -L -o solver.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/host/solver.py
11.3 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """DPM-Solver++(2M) with denoise-to-zero, as Autoware runs it around the decoder (multi_step mode). | |
| Port of ``PKG/src/inference/solver/dpm_solver.cpp`` (guidance disabled, Autoware's default and decision D11) and of | |
| the prefix constraint of ``PKG/src/inference/multi_step_inference.cpp:300-340``. The scalar schedule is float32 math | |
| evaluated with the C library's ``expf`` / ``logf`` / ``expm1f`` / ``sqrtf`` (what ``std::exp(float)`` etc. call in | |
| the node) in the C++ operation order, so the timesteps and update coefficients are the node's bit for bit; numpy's | |
| own float32 transcendentals differ from glibc in the last bit for many arguments. Without a loadable libm the module | |
| falls back to numpy (``SCALAR_MATH`` says which). | |
| For ``steps = 10``: 11 decoder evaluations at ``timesteps[0..9]`` and ``1/N = 0.001`` (denoise-to-zero); the 11 | |
| published iterates (``denoising_steps``) are the corrected ``x`` after each update and after the final evaluation. | |
| The update formulas, with ``x`` and the model outputs float32 arrays and every scalar a float32:: | |
| first (step 1): x = (sigma_t / sigma_s) * x - (alpha_t * phi_1) * m | |
| second (steps 2..): d = (m0 - m1) / r0 | |
| x = (sigma_t / sigma_0) * x - (alpha_t * phi_1) * m0 - (0.5 * (alpha_t * phi_1)) * d | |
| :class:`SolverPlan` lists these scalars per update (the RT-dev coefficient tables of the on-device loop). | |
| """ | |
| from __future__ import annotations | |
| import ctypes | |
| import ctypes.util | |
| from dataclasses import dataclass, field | |
| from typing import Callable, List, Optional, Tuple | |
| import numpy as np | |
| from ..reference import config as C | |
| __all__ = ["SCALAR_MATH", "marginal_log_mean_coeff", "marginal_alpha", "marginal_std", "marginal_lambda", | |
| "inverse_lambda", "log_snr_timesteps", "SolverUpdate", "SolverPlan", "solver_plan", "first_update", | |
| "second_update", "dpm_solver_sample", "apply_prefix_constraint", "SampleResult"] | |
| f32 = np.float32 | |
| _T, _N = C.NOISE_SCHEDULE_T, C.NOISE_SCHEDULE_TOTAL_N | |
| _B0, _B1 = C.NOISE_SCHEDULE_BETA0, C.NOISE_SCHEDULE_BETA1 | |
| class _ScalarMath: | |
| """float32 ``exp`` / ``log`` / ``expm1`` / ``sqrt`` of the C library, loaded on first use (numpy fallback).""" | |
| def __init__(self) -> None: | |
| self._fns = None | |
| self.source = "unloaded" | |
| def _load(self) -> None: | |
| fns = None | |
| try: | |
| name = ctypes.util.find_library("m") or "libm.so.6" | |
| libm = ctypes.CDLL(name) | |
| fns = {} | |
| for fn in ("expf", "logf", "expm1f", "sqrtf"): | |
| f = getattr(libm, fn) | |
| f.restype, f.argtypes = ctypes.c_float, [ctypes.c_float] | |
| fns[fn] = f | |
| self.source = f"libm ({name})" | |
| except (OSError, AttributeError): | |
| fns = None | |
| self.source = "numpy (no C math library found)" | |
| self._fns = fns or {} | |
| def _call(self, fn: str, x: float, fallback: Callable) -> np.float32: | |
| if self._fns is None: | |
| self._load() | |
| if fn in self._fns: | |
| return f32(self._fns[fn](float(f32(x)))) | |
| return f32(fallback(f32(x))) | |
| def exp(self, x): | |
| return self._call("expf", x, np.exp) | |
| def log(self, x): | |
| return self._call("logf", x, np.log) | |
| def expm1(self, x): | |
| return self._call("expm1f", x, np.expm1) | |
| def sqrt(self, x): | |
| return self._call("sqrtf", x, np.sqrt) | |
| SCALAR_MATH = _ScalarMath() | |
| _m = SCALAR_MATH | |
| # ---- noise schedule (dpm_solver.cpp:49-94), float32 in the C++ operation order --------------------------------------- | |
| def marginal_log_mean_coeff(t) -> np.float32: | |
| """``-0.25f * t * t * (beta1 - beta0) - 0.5f * t * beta0``.""" | |
| t = f32(t) | |
| return f32(f32(f32(f32(-0.25) * t) * t) * f32(_B1 - _B0)) - f32(f32(f32(0.5) * t) * _B0) | |
| def marginal_alpha(t) -> np.float32: | |
| return _m.exp(marginal_log_mean_coeff(t)) | |
| def marginal_std(t) -> np.float32: | |
| return _m.sqrt(f32(f32(1.0) - _m.exp(f32(f32(2.0) * marginal_log_mean_coeff(t))))) | |
| def marginal_lambda(t) -> np.float32: | |
| lmc = marginal_log_mean_coeff(t) | |
| log_std = f32(f32(0.5) * _m.log(f32(f32(1.0) - _m.exp(f32(f32(2.0) * lmc))))) | |
| return f32(lmc - log_std) | |
| def _log_add_exp(a, b) -> np.float32: | |
| m = max(f32(a), f32(b)) | |
| return f32(m + _m.log(f32(_m.exp(f32(f32(a) - m)) + _m.exp(f32(f32(b) - m))))) | |
| def inverse_lambda(lam) -> np.float32: | |
| beta_delta = f32(_B1 - _B0) | |
| tmp = f32(f32(f32(2.0) * beta_delta) * _log_add_exp(f32(f32(-2.0) * f32(lam)), f32(0.0))) | |
| delta = f32(f32(_B0 * _B0) + tmp) | |
| return f32(f32(tmp / f32(_m.sqrt(delta) + _B0)) / beta_delta) | |
| def log_snr_timesteps(steps: int) -> List[np.float32]: | |
| """``steps + 1`` times uniform in log-SNR between t = 1 and t = 1/N (``dpm_solver.cpp:80-94``).""" | |
| t0 = f32(f32(1.0) / _N) | |
| lam_t, lam_0 = marginal_lambda(_T), marginal_lambda(t0) | |
| out = [] | |
| for i in range(steps + 1): | |
| ratio = f32(f32(i) / f32(steps)) | |
| out.append(inverse_lambda(f32(lam_t + f32(f32(lam_0 - lam_t) * ratio)))) | |
| return out | |
| # ---- updates (dpm_solver.cpp:96-136) ------------------------------------------------------------------------------ | |
| def first_update(x_s: np.ndarray, model_s: np.ndarray, s, t) -> np.ndarray: | |
| h = f32(marginal_lambda(t) - marginal_lambda(s)) | |
| sigma_s, sigma_t, alpha_t = marginal_std(s), marginal_std(t), marginal_alpha(t) | |
| phi_1 = _m.expm1(f32(-h)) | |
| a, b = f32(sigma_t / sigma_s), f32(alpha_t * phi_1) | |
| return (a * x_s.astype(np.float32) - b * model_s.astype(np.float32)).astype(np.float32) | |
| def second_update(x_s: np.ndarray, model_prev: Tuple[np.ndarray, np.ndarray], t_prev: Tuple[float, float], | |
| t) -> np.ndarray: | |
| m1, m0 = model_prev # model_prev_list[0] (older), model_prev_list[1] (newer) | |
| t1, t0 = t_prev | |
| lam1, lam0, lam_t = marginal_lambda(t1), marginal_lambda(t0), marginal_lambda(t) | |
| sigma0, sigma_t, alpha_t = marginal_std(t0), marginal_std(t), marginal_alpha(t) | |
| h0, h = f32(lam0 - lam1), f32(lam_t - lam0) | |
| r0 = f32(h0 / h) | |
| phi_1 = _m.expm1(f32(-h)) | |
| a, b = f32(sigma_t / sigma0), f32(alpha_t * phi_1) | |
| c = f32(f32(0.5) * b) | |
| d1_0 = ((m0.astype(np.float32) - m1.astype(np.float32)) / r0).astype(np.float32) | |
| return (a * x_s.astype(np.float32) - b * m0.astype(np.float32) - c * d1_0).astype(np.float32) | |
| class SolverUpdate: | |
| """One solver update to time ``t``: ``x = a * x - b * m0 - c * (m0 - m1) / r0`` (``c = 0`` for the first-order | |
| update, which has no ``m1``).""" | |
| order: int | |
| t: float | |
| a: float | |
| b: float | |
| c: float | |
| r0: float | |
| class SolverPlan: | |
| """Everything the on-device loop needs for ``steps`` (the RT-dev tables of PLAN.md 2.12): the decoder evaluation | |
| times (``eval_times[k]`` feeds evaluation ``k``; the t-embedding / adaLN tables are built for these values), the | |
| update after each of the first ``steps`` evaluations, and the published iterate times.""" | |
| steps: int | |
| timesteps: Tuple[float, ...] | |
| eval_times: Tuple[float, ...] | |
| updates: Tuple[SolverUpdate, ...] | |
| denoising_timesteps: Tuple[float, ...] | |
| scalar_math: str = field(default="") | |
| def solver_plan(steps: int = C.DPM_SOLVER_STEPS) -> SolverPlan: | |
| if steps < C.DPM_SOLVER_ORDER: | |
| raise ValueError("DpmSolver steps must be greater than or equal to solver order.") | |
| ts = log_snr_timesteps(steps) | |
| updates = [] | |
| for step in range(1, steps + 1): | |
| t = ts[step] | |
| if step < C.DPM_SOLVER_ORDER: | |
| s = _T if step == 1 else ts[step - 1] | |
| h = f32(marginal_lambda(t) - marginal_lambda(s)) | |
| a = f32(marginal_std(t) / marginal_std(s)) | |
| b = f32(marginal_alpha(t) * _m.expm1(f32(-h))) | |
| updates.append(SolverUpdate(1, float(t), float(a), float(b), 0.0, 1.0)) | |
| else: | |
| t1 = _T if step - 2 == 0 else ts[step - 2] | |
| t0 = ts[step - 1] | |
| lam1, lam0, lam_t = marginal_lambda(t1), marginal_lambda(t0), marginal_lambda(t) | |
| h0, h = f32(lam0 - lam1), f32(lam_t - lam0) | |
| b = f32(marginal_alpha(t) * _m.expm1(f32(-h))) | |
| updates.append(SolverUpdate(2, float(t), float(f32(marginal_std(t) / marginal_std(t0))), float(b), | |
| float(f32(f32(0.5) * b)), float(f32(h0 / h)))) | |
| final_t = f32(f32(1.0) / _N) | |
| return SolverPlan(steps, tuple(float(t) for t in ts), tuple(float(t) for t in ts[:steps]) + (float(final_t),), | |
| tuple(updates), tuple(float(t) for t in ts[1:]) + (float(final_t),), SCALAR_MATH.source) | |
| # ---- the loop (dpm_solver.cpp:138-234) --------------------------------------------------------------------------- | |
| class SampleResult: | |
| final_x: np.ndarray | |
| denoising_steps: List[np.ndarray] | |
| denoising_timesteps: List[float] | |
| eval_times: List[float] | |
| nfe: int | |
| def apply_prefix_constraint(x: np.ndarray, current_states: np.ndarray) -> np.ndarray: | |
| """``x[agent, 0, :] = current_states[agent]`` in place (``multi_step_inference.cpp:327-340``); ``x`` is | |
| ``[321, 81, 4]`` (or with a leading batch dim).""" | |
| x[..., 0, :] = current_states | |
| return x | |
| def dpm_solver_sample(initial_x: np.ndarray, model_fn: Callable[[np.ndarray, np.float32], np.ndarray], | |
| correcting_fn: Callable[[np.ndarray], None], steps: int = C.DPM_SOLVER_STEPS, | |
| on_iterate: Optional[Callable[[int, np.ndarray], None]] = None) -> SampleResult: | |
| """``DpmSolver::sample`` without guidance. ``model_fn(x, t)`` is one decoder evaluation (x0 prediction), | |
| ``correcting_fn(x)`` the in-place prefix constraint.""" | |
| if steps < C.DPM_SOLVER_ORDER: | |
| raise ValueError("DpmSolver steps must be greater than or equal to solver order.") | |
| x = np.array(initial_x, dtype=np.float32, copy=True) | |
| correcting_fn(x) | |
| ts = log_snr_timesteps(steps) | |
| steps_out: List[np.ndarray] = [] | |
| step_ts: List[float] = [] | |
| eval_times: List[float] = [] | |
| def evaluate(xx: np.ndarray, t) -> np.ndarray: | |
| eval_times.append(float(t)) | |
| return np.asarray(model_fn(xx, f32(t)), np.float32) | |
| def record(xx: np.ndarray, t) -> None: | |
| steps_out.append(xx.copy()) | |
| step_ts.append(float(t)) | |
| if on_iterate is not None: | |
| on_iterate(len(steps_out) - 1, xx) | |
| t_prev = [_T] | |
| m_prev = [evaluate(x, ts[0])] | |
| for step in range(1, C.DPM_SOLVER_ORDER): | |
| t = ts[step] | |
| x = first_update(x, m_prev[-1], t_prev[-1], t) | |
| correcting_fn(x) | |
| record(x, t) | |
| t_prev.append(t) | |
| m_prev.append(evaluate(x, t)) | |
| for step in range(C.DPM_SOLVER_ORDER, steps + 1): | |
| t = ts[step] | |
| x = second_update(x, (m_prev[0], m_prev[1]), (t_prev[0], t_prev[1]), t) | |
| correcting_fn(x) | |
| record(x, t) | |
| t_prev = [t_prev[1], t] | |
| m_prev = [m_prev[1], None] | |
| if step < steps: | |
| m_prev[1] = evaluate(x, t) | |
| final_t = f32(f32(1.0) / _N) | |
| x = evaluate(x, final_t) # denoise to zero: the model output itself, no guidance on the last step | |
| x = np.array(x, dtype=np.float32, copy=True) | |
| correcting_fn(x) | |
| record(x, final_t) | |
| return SampleResult(x, steps_out, step_ts, eval_times, len(eval_times)) | |