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