Download code/tt_diffusion_planner/tests/test_host_solver.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.15 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_host_solver.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tests/test_host_solver.py
-
curl -L -o test_host_solver.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_host_solver.py
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) | |
| 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) | |
| 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) | |