changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
8.15 kB
# SPDX-License-Identifier: Apache-2.0
"""DPM-Solver++(2M) host code (no device, no weights): the schedule of ``dpm_solver.cpp`` bit for bit, the loop
against the verified research port (``dp_common.dpm_solver_sample``) and against an independent float64 transcription
of the update formulas, the prefix constraint, and the coefficient plan the on-device loop will use.
TT_VISIBLE_DEVICES=none python -m pytest -q code/tt_diffusion_planner/tests/test_host_solver.py
"""
from __future__ import annotations
import math
import numpy as np
import pytest
from tt_diffusion_planner.host import solver as S
from tt_diffusion_planner.reference import config as C
from tt_diffusion_planner.tests import _research as R
def test_timesteps_steps10_match_the_spec():
ts = S.log_snr_timesteps(10)
assert len(ts) == 11 and all(isinstance(t, np.float32) for t in ts)
np.testing.assert_allclose(np.asarray(ts, np.float64), C.DPM_TIMESTEPS_STEPS10, rtol=0, atol=6e-9)
assert ts[0] == np.float32(1.0) and ts[-1] == np.float32(0.0010006181)
assert np.all(np.diff(np.asarray(ts)) < 0)
def test_schedule_against_float64_formulas():
"""The float32 schedule follows the exact VP-linear schedule up to float32 rounding; near t = 0 the C++ float
expression ``1 - exp(2 lmc)`` cancels (lmc ~ -5.5e-5 at t = 1e-3), which costs ~1e-4 relative in sigma and is
reproduced on purpose."""
b0, b1 = 0.1, 20.0
for t in (1.0, 0.5, 0.05, 0.001):
tol = 1e-5 if t > 0.01 else 3e-4
lmc = -0.25 * t * t * (b1 - b0) - 0.5 * t * b0
assert math.isclose(S.marginal_log_mean_coeff(t), lmc, rel_tol=1e-6)
assert math.isclose(S.marginal_alpha(t), math.exp(lmc), rel_tol=1e-6)
assert math.isclose(S.marginal_std(t), math.sqrt(1 - math.exp(2 * lmc)), rel_tol=tol)
lam = lmc - 0.5 * math.log(1 - math.exp(2 * lmc))
assert math.isclose(S.marginal_lambda(t), lam, rel_tol=tol)
assert math.isclose(S.inverse_lambda(np.float32(lam)), t, rel_tol=10 * tol)
@pytest.mark.skipif(R.load_script("dp_common") is None, reason="research scripts not present")
def test_timesteps_equal_research_port(monkeypatch):
"""Same operation order as ``dp_common.py``: with numpy's float32 transcendentals (the research port's choice)
every timestep is bit-identical; with the C library (the node's) the deployed steps = 10 still are, other step
counts may move by one ulp."""
dp = R.load_script("dp_common")
np.testing.assert_array_equal(np.asarray(S.log_snr_timesteps(10)), np.asarray(dp.log_snr_timesteps(10)))
for steps in (2, 5, 20):
np.testing.assert_allclose(np.asarray(S.log_snr_timesteps(steps)), np.asarray(dp.log_snr_timesteps(steps)),
rtol=3e-7, atol=0)
monkeypatch.setattr(S.SCALAR_MATH, "_fns", {})
for steps in (2, 5, 10, 20):
np.testing.assert_array_equal(np.asarray(S.log_snr_timesteps(steps)), np.asarray(dp.log_snr_timesteps(steps)))
def _toy_model(x, t):
"""A deterministic x0-predictor stand-in: mixes x with a time-dependent target."""
target = np.sin(np.arange(x.size, dtype=np.float32).reshape(x.shape) * np.float32(0.01))
return (np.float32(0.3) * x + np.float32(1.0 - 0.3) * target * np.float32(1.0 + t)).astype(np.float32)
def _float64_dpm(x, steps, model, correct):
"""Independent float64 transcription of DPM-Solver++(2M) + denoise-to-zero (dpm_solver.cpp:96-234)."""
b0, b1 = 0.1, 20.0
lmc = lambda t: -0.25 * t * t * (b1 - b0) - 0.5 * t * b0 # noqa: E731
alpha = lambda t: math.exp(lmc(t)) # noqa: E731
sigma = lambda t: math.sqrt(1 - math.exp(2 * lmc(t))) # noqa: E731
lam = lambda t: math.log(alpha(t)) - math.log(sigma(t)) # noqa: E731
ts = [float(t) for t in S.log_snr_timesteps(steps)]
x = x.astype(np.float64)
correct(x)
ms, tp = [model(x, ts[0])], [1.0]
t = ts[1]
h = lam(t) - lam(tp[-1])
x = sigma(t) / sigma(tp[-1]) * x - alpha(t) * math.expm1(-h) * ms[-1]
correct(x)
tp.append(t)
ms.append(model(x, t))
for step in range(2, steps + 1):
t = ts[step]
h0, h = lam(tp[1]) - lam(tp[0]), lam(t) - lam(tp[1])
phi = math.expm1(-h)
d = (ms[1] - ms[0]) / (h0 / h)
x = sigma(t) / sigma(tp[1]) * x - alpha(t) * phi * ms[1] - 0.5 * alpha(t) * phi * d
correct(x)
tp = [tp[1], t]
if step < steps:
ms = [ms[1], model(x, t)]
x = model(x, 0.001)
correct(x)
return x
def test_loop_against_float64_transcription():
rng = np.random.default_rng(1)
x0 = rng.normal(size=(321, 81, 4)).astype(np.float32)
cs = rng.normal(size=(321, 4)).astype(np.float32)
calls = []
def model(x, t):
calls.append(float(t))
return _toy_model(np.asarray(x, np.float32), float(t))
res = S.dpm_solver_sample(x0, model, lambda x: S.apply_prefix_constraint(x, cs), 10)
assert res.nfe == 11 == len(calls)
plan = S.solver_plan(10)
np.testing.assert_allclose(calls, plan.eval_times, rtol=0, atol=0)
assert len(res.denoising_steps) == 11 and res.denoising_timesteps == list(plan.denoising_timesteps)
for x in res.denoising_steps:
np.testing.assert_array_equal(x[:, 0, :], cs)
want = _float64_dpm(x0, 10, lambda x, t: _toy_model(x.astype(np.float32), t).astype(np.float64),
lambda x: S.apply_prefix_constraint(x, cs))
np.testing.assert_allclose(res.final_x, want, rtol=0, atol=2e-5)
@pytest.mark.skipif(R.load_script("dp_common") is None, reason="research scripts not present")
def test_loop_against_research_port(monkeypatch):
"""``dp_common.py`` evaluates the scalars with numpy instead of the C library: with the same scalar math the loop
is bit-identical (same update order), with the C library it differs by the scalars' last bits only."""
dp = R.load_script("dp_common")
rng = np.random.default_rng(2)
x0 = rng.normal(size=(1, 321, 81, 4)).astype(np.float32)
cs = rng.normal(size=(1, 321, 4)).astype(np.float32)
b_final, b_steps, b_ts, b_nfe = dp.dpm_solver_sample(x0, _toy_model, dp.make_prefix_constraint(cs[0]), 10)
a = S.dpm_solver_sample(x0, _toy_model, lambda x: S.apply_prefix_constraint(x, cs), 10)
assert b_nfe == a.nfe == 11
np.testing.assert_array_equal(np.asarray(a.denoising_timesteps, np.float32), np.asarray(b_ts, np.float32))
np.testing.assert_allclose(a.final_x, b_final, rtol=0, atol=2e-5)
monkeypatch.setattr(S.SCALAR_MATH, "_fns", {})
a = S.dpm_solver_sample(x0, _toy_model, lambda x: S.apply_prefix_constraint(x, cs), 10)
np.testing.assert_array_equal(a.final_x, b_final)
for p, q in zip(a.denoising_steps, b_steps):
np.testing.assert_array_equal(p, q)
def test_plan_reproduces_the_updates_bit_for_bit():
"""``x = a x - b m0 - c (m0 - m1) / r0`` with the plan's float32 scalars equals first_update / second_update."""
rng = np.random.default_rng(3)
plan = S.solver_plan(10)
ts = S.log_snr_timesteps(10)
assert [u.order for u in plan.updates] == [1] + [2] * 9
x, m0, m1 = (rng.normal(size=(321, 81, 4)).astype(np.float32) for _ in range(3))
u = plan.updates[0]
want = S.first_update(x, m0, C.NOISE_SCHEDULE_T, ts[1])
got = (np.float32(u.a) * x - np.float32(u.b) * m0).astype(np.float32)
np.testing.assert_array_equal(got, want)
for k in range(2, 11):
u = plan.updates[k - 1]
t1 = C.NOISE_SCHEDULE_T if k == 2 else ts[k - 2]
want = S.second_update(x, (m1, m0), (t1, ts[k - 1]), ts[k])
d = ((m0 - m1) / np.float32(u.r0)).astype(np.float32)
got = (np.float32(u.a) * x - np.float32(u.b) * m0 - np.float32(u.c) * d).astype(np.float32)
np.testing.assert_array_equal(got, want, err_msg=f"update {k}")
def test_scalar_math_is_the_c_library():
S.marginal_alpha(0.5)
assert S.SCALAR_MATH.source.startswith("libm"), S.SCALAR_MATH.source
def test_steps_below_order_raise():
with pytest.raises(ValueError):
S.solver_plan(1)
with pytest.raises(ValueError):
S.dpm_solver_sample(np.zeros((321, 81, 4), np.float32), _toy_model, lambda x: None, 1)