# 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)