From d0941eff11a8cf53f1925bc0bc7c6d0f10fdaa2f Mon Sep 17 00:00:00 2001 From: gertl Date: Wed, 5 Aug 2026 19:21:44 +0200 Subject: [PATCH 1/2] Add DPM-Solver++(2M) diffusion sampler DPMSolverPlusPlus2M reaches second order with one denoiser evaluation per step by reusing the previous step's data prediction, where HeunSolver needs two. On a closed-form probability-flow ODE it is roughly 4x more accurate than Heun at matched network evaluations. The update is derived for the EDM parameterization (sigma(t) = t, alpha(t) = 1), so sample() rejects schedulers that do not declare the EDM parameterization. Compatibility is declared by a new is_edm_parameterization attribute rather than probed numerically, since resolving a tensor comparison to a Python bool would break fullgraph tracing. The solver is stateful, so sample() clears its history before each trajectory and in a finally block, gated on a _requires_state_reset marker so that a pre-existing solver with its own reset() is untouched. Neither the Solver nor the NoiseScheduler protocol gains a member: both are runtime checkable, so that would flip isinstance() for existing user-defined implementations. step() avoids data-dependent Python branches, which would force a host sync every step and break fullgraph tracing. Step sizes are computed in at least float32, the step following a repeated timestep falls back to first order, and a state already at t = 0 is returned unchanged. Tests cover the non-uniform-step coefficients, numerical edge cases, compilation, DPS guidance and the sampler lifecycle; the registry key is "dpmpp_2m". Co-Authored-By: Claude Opus 5 Signed-off-by: gertl --- CHANGELOG.md | 4 + docs/api/diffusion/samplers.rst | 13 + .../noise_schedulers/noise_schedulers.py | 24 ++ physicsnemo/diffusion/samplers/__init__.py | 1 + physicsnemo/diffusion/samplers/samplers.py | 123 +++++-- physicsnemo/diffusion/samplers/solvers.py | 230 ++++++++++++++ .../data/test_solvers_dpmpp_2m_1d_step.pth | Bin 0 -> 2243 bytes .../data/test_solvers_dpmpp_2m_2d_step.pth | Bin 0 -> 3011 bytes .../data/test_solvers_dpmpp_2m_3d_step.pth | Bin 0 -> 2883 bytes test/diffusion/test_samplers.py | 293 +++++++++++++++++ test/diffusion/test_solvers.py | 300 ++++++++++++++++++ 11 files changed, 964 insertions(+), 24 deletions(-) create mode 100644 test/diffusion/data/test_solvers_dpmpp_2m_1d_step.pth create mode 100644 test/diffusion/data/test_solvers_dpmpp_2m_2d_step.pth create mode 100644 test/diffusion/data/test_solvers_dpmpp_2m_3d_step.pth diff --git a/CHANGELOG.md b/CHANGELOG.md index 6d796875a5..fb1671ac3d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -213,6 +213,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 next event (the new particle's features and inter-event delay) from the current particle population, an optional background mesh, and the simulation time. Independent rollouts form an ensemble for uncertainty quantification. +- Adds `DPMSolverPlusPlus2M` (string key `"dpmpp_2m"`) to + `physicsnemo.diffusion.samplers`, a second-order multistep DPM-Solver++ that + reuses the previous step's data prediction and so needs only one denoiser + evaluation per step. Requires a noise scheduler in the EDM parameterization. ### Changed diff --git a/docs/api/diffusion/samplers.rst b/docs/api/diffusion/samplers.rst index 5b06ea5c58..b130c3e248 100644 --- a/docs/api/diffusion/samplers.rst +++ b/docs/api/diffusion/samplers.rst @@ -406,6 +406,11 @@ There are two ways to use solvers: - ``"edm_stochastic_heun"`` --- :class:`~physicsnemo.diffusion.samplers.solvers.EDMStochasticHeunSolver`. Second-order with configurable stochastic noise injection. +- ``"dpmpp_2m"`` --- + :class:`~physicsnemo.diffusion.samplers.solvers.DPMSolverPlusPlus2M`. + Second-order multistep, one denoiser evaluation per step. Requires a noise + scheduler in the EDM parameterization (:math:`\sigma(t) = t`, + :math:`\alpha(t) = 1`). **Custom solvers** can be defined by implementing the :class:`~physicsnemo.diffusion.samplers.solvers.Solver` protocol: any object @@ -502,6 +507,14 @@ Solvers :members: :exclude-members: __init__ +:code:`DPMSolverPlusPlus2M` +^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +.. autoclass:: physicsnemo.diffusion.samplers.solvers.DPMSolverPlusPlus2M + :show-inheritance: + :members: + :exclude-members: __init__ + Guidance ~~~~~~~~ diff --git a/physicsnemo/diffusion/noise_schedulers/noise_schedulers.py b/physicsnemo/diffusion/noise_schedulers/noise_schedulers.py index 311d163507..f817cb8ec5 100644 --- a/physicsnemo/diffusion/noise_schedulers/noise_schedulers.py +++ b/physicsnemo/diffusion/noise_schedulers/noise_schedulers.py @@ -324,6 +324,16 @@ class LinearGaussianNoiseScheduler(ABC, NoiseScheduler): - :meth:`init_latents`: Initialize latent state (sampling) - :meth:`get_denoiser`: Get ODE/SDE RHS (sampling) + Attributes + ---------- + is_edm_parameterization : bool, default=False + Whether the schedule satisfies :math:`\sigma(t) = t` and + :math:`\alpha(t) = 1`, so that the diffusion time is the noise level. + Solvers derived only for this case, such as + :class:`~physicsnemo.diffusion.samplers.solvers.DPMSolverPlusPlus2M`, + are rejected by :func:`~physicsnemo.diffusion.samplers.sample` + otherwise. Subclasses satisfying both identities may set it to ``True``. + Examples -------- **Example 1:** A minimal EDM-like noise schedule. Only the abstract methods @@ -335,6 +345,10 @@ class LinearGaussianNoiseScheduler(ABC, NoiseScheduler): ... ) >>> >>> class SimpleEDMScheduler(LinearGaussianNoiseScheduler): + ... # sigma(t) = t and alpha(t) = 1 below, so opt in to the solvers + ... # that require the EDM parameterization. + ... is_edm_parameterization = True + ... ... def __init__(self, sigma_min=0.002, sigma_max=80.0, rho=7.0): ... self.sigma_min = sigma_min ... self.sigma_max = sigma_max @@ -386,6 +400,10 @@ class LinearGaussianNoiseScheduler(ABC, NoiseScheduler): """ + # Kept off the NoiseScheduler protocol, which is runtime_checkable: a new + # required member would break isinstance for duck-typed schedulers. + is_edm_parameterization: bool = False + @abstractmethod def sigma( self, @@ -1302,6 +1320,8 @@ class EDMNoiseScheduler(LinearGaussianNoiseScheduler): torch.Size([4, 3]) """ + is_edm_parameterization: bool = True + def __init__( self, sigma_min: float = 0.002, @@ -1810,6 +1830,8 @@ class IDDPMNoiseScheduler(LinearGaussianNoiseScheduler): torch.Size([4, 3, 8, 8]) """ + is_edm_parameterization: bool = True + def __init__( self, sigma_min: float = 0.002, @@ -2284,6 +2306,8 @@ class StudentTEDMNoiseScheduler(LinearGaussianNoiseScheduler): torch.Size([4, 3, 8, 8]) """ + is_edm_parameterization: bool = True + def __init__( self, sigma_min: float = 0.002, diff --git a/physicsnemo/diffusion/samplers/__init__.py b/physicsnemo/diffusion/samplers/__init__.py index 2ef817e33f..76d0806e28 100644 --- a/physicsnemo/diffusion/samplers/__init__.py +++ b/physicsnemo/diffusion/samplers/__init__.py @@ -18,6 +18,7 @@ from .legacy_stochastic_sampler import stochastic_sampler # noqa: F401 from .samplers import sample # noqa: F401 from .solvers import ( # noqa: F401 + DPMSolverPlusPlus2M, EDMStochasticEulerSolver, EDMStochasticHeunSolver, EulerSolver, diff --git a/physicsnemo/diffusion/samplers/samplers.py b/physicsnemo/diffusion/samplers/samplers.py index 5432bafb50..c709218b53 100644 --- a/physicsnemo/diffusion/samplers/samplers.py +++ b/physicsnemo/diffusion/samplers/samplers.py @@ -29,6 +29,7 @@ from physicsnemo.domain_parallel.shard_tensor import scatter_tensor from .solvers import ( + DPMSolverPlusPlus2M, EDMStochasticEulerSolver, EDMStochasticHeunSolver, EulerSolver, @@ -41,9 +42,49 @@ "heun": HeunSolver, "edm_stochastic_euler": EDMStochasticEulerSolver, "edm_stochastic_heun": EDMStochasticHeunSolver, + "dpmpp_2m": DPMSolverPlusPlus2M, } +def _check_edm_parameterization( + noise_scheduler: NoiseScheduler, solver: Solver +) -> None: + r"""Raise if ``noise_scheduler`` is not in the EDM parameterization. + + Compatibility is read from the scheduler's ``is_edm_parameterization`` + attribute rather than probed numerically, because converting a tensor + comparison to a Python boolean would break + ``torch.compile(..., fullgraph=True)``. + + Parameters + ---------- + noise_scheduler : NoiseScheduler + The scheduler to validate. A + :class:`~physicsnemo.diffusion.noise_schedulers.DomainParallelNoiseScheduler` + is unwrapped and its inner scheduler is validated instead. + solver : Solver + The solver requiring the parameterization; used in the error message. + + Returns + ------- + None + + Raises + ------ + ValueError + If the scheduler is not declared to be in the EDM parameterization. + """ + # Unwrap the domain-parallel adapter: it only changes tensor placement. + inner = getattr(noise_scheduler, "inner_scheduler", noise_scheduler) + + if not getattr(inner, "is_edm_parameterization", False): + raise ValueError( + f"{type(solver).__name__} requires a noise scheduler in the EDM " + f"parameterization (sigma(t) = t, alpha(t) = 1), but " + f"{type(inner).__name__} does not set is_edm_parameterization=True." + ) + + def _maybe_replicate_timesteps( t_steps: Float[Tensor, " N_plus_1"], xN: Float[Tensor, " B *dims"], @@ -78,7 +119,9 @@ def sample( xN: Float[Tensor, " B *dims"], noise_scheduler: NoiseScheduler, num_steps: int, - solver: Literal["euler", "heun", "edm_stochastic_euler", "edm_stochastic_heun"] + solver: Literal[ + "euler", "heun", "edm_stochastic_euler", "edm_stochastic_heun", "dpmpp_2m" + ] | Solver = "heun", time_steps: Float[Tensor, " N_plus_1"] | None = None, solver_options: Dict[str, Any] | None = None, @@ -216,6 +259,11 @@ def denoiser( the EDM paper with configurable noise injection. See :class:`~physicsnemo.diffusion.samplers.solvers.EDMStochasticHeunSolver`. + * ``"dpmpp_2m"``: Second-order multistep DPM-Solver++, using a single + denoiser evaluation per step. Requires a noise scheduler in the EDM + parameterization. + See :class:`~physicsnemo.diffusion.samplers.solvers.DPMSolverPlusPlus2M`. + time_steps : Tensor | None, default=None Optional 1D tensor of shape :math:`(N + 1,)` containing explicit diffusion time values :math:`t_N, t_{N-1}, ..., t_0` in decreasing @@ -320,6 +368,11 @@ def denoiser( >>> >>> # Define a minimal EDM-like scheduler from scratch >>> class MinimalScheduler: + ... # sigma=t and alpha=1 below, so opt in to the solvers that require + ... # the EDM parameterization. Duck-typed schedulers can set this too: + ... # it is read defensively, not required by the protocol. + ... is_edm_parameterization = True + ... ... def timesteps(self, num_steps, *, device=None, dtype=None): ... return torch.linspace(1.0, 0.0, num_steps + 1, ... device=device, dtype=dtype) @@ -391,6 +444,23 @@ def denoiser( # callers that intentionally backprop through sample() are unaffected. outer_grad_enabled = torch.is_grad_enabled() + # Reject incompatible solver/scheduler parameterizations before evaluating + # the denoiser. + if getattr(solver_, "_requires_edm_parameterization", False): + _check_edm_parameterization(noise_scheduler, solver_) + + # Stateful solvers opt into a reset hook, cleared here so a reused instance + # cannot leak state from a previous trajectory. Gated on the marker rather + # than on the presence of a ``reset`` method: ``reset`` is a common name, and + # a pre-existing solver that happens to define one for its own purposes must + # keep its previous behavior. The hook stays outside the runtime-checkable + # ``Solver`` protocol so that solvers implementing only ``step`` still match. + reset = getattr(solver_, "reset", None) + if not getattr(solver_, "_requires_state_reset", False): + reset = None + if callable(reset): + reset() + # Main sampling loop samples: List[Tensor] = [] x = xN @@ -404,26 +474,31 @@ def denoiser( f"valid indices are in range(0, {n_steps})." ) - for i in range(n_steps): - t_cur = t_steps[i] - t_next = t_steps[i + 1] - - # Expand t to batch dimension: scalar -> (B,) - batch_size = x.shape[0] - t_cur_batch = t_cur.expand(batch_size) - t_next_batch = t_next.expand(batch_size) - - # Perform one solver step - x = solver_.step(x, t_cur_batch, t_next_batch) - if not outer_grad_enabled: - x = x.detach() - - # Collect sample if requested - if time_eval is not None and i in time_eval: - samples.append(x.clone()) - - # Return based on time_eval - if time_eval is not None: - return samples - - return x + try: + for i in range(n_steps): + t_cur = t_steps[i] + t_next = t_steps[i + 1] + + # Expand t to batch dimension: scalar -> (B,) + batch_size = x.shape[0] + t_cur_batch = t_cur.expand(batch_size) + t_next_batch = t_next.expand(batch_size) + + # Perform one solver step + x = solver_.step(x, t_cur_batch, t_next_batch) + if not outer_grad_enabled: + x = x.detach() + + # Collect sample if requested + if time_eval is not None and i in time_eval: + samples.append(x.clone()) + + # Return based on time_eval + if time_eval is not None: + return samples + + return x + finally: + # Release cached solver state on success or failure. + if callable(reset): + reset() diff --git a/physicsnemo/diffusion/samplers/solvers.py b/physicsnemo/diffusion/samplers/solvers.py index d663a58d2f..25409c13cf 100644 --- a/physicsnemo/diffusion/samplers/solvers.py +++ b/physicsnemo/diffusion/samplers/solvers.py @@ -792,3 +792,233 @@ def step( x_next = mask_bc * x_heun + (1 - mask_bc) * x_euler return x_next + + +class DPMSolverPlusPlus2M(Solver): + r""" + DPM-Solver++(2M): second-order multistep solver for diffusion ODEs. + + Unlike :class:`HeunSolver`, which attains second order with *two* denoiser + evaluations per step, this solver attains second order with a *single* + evaluation per step by reusing the previous step's data prediction + (a linear-multistep, Adams--Bashforth-style scheme applied to the + probability-flow ODE in exponential-integrator form). For a fixed budget of + network function evaluations it can therefore take twice as many steps as a + second-order single-step method. + + This implementation targets the variance-exploding (EDM) parameterization, + where the diffusion time equals the noise level (:math:`t = \sigma`) and + :math:`\alpha_t = 1`. Writing :math:`\lambda = -\ln \sigma`, with step size + :math:`h = \lambda_{n-1} - \lambda_n` so that + :math:`e^{-h} = \sigma_{n-1} / \sigma_n`, the update is + + .. math:: + \mathbf{x}_{n-1} = e^{-h}\, \mathbf{x}_n + + \left(1 - e^{-h}\right) \bar{\mathbf{D}}_n , + + where the extrapolated data prediction is + + .. math:: + \bar{\mathbf{D}}_n = \left(1 + \frac{1}{2r}\right) \mathbf{D}_n + - \frac{1}{2r} \mathbf{D}_{n+1} , + \qquad r = \frac{h_{n+1}}{h_n} , + + and :math:`\mathbf{D}_n` is the denoised (data-space) prediction at step + :math:`n`. On the first step -- and whenever no history is available -- the + scheme falls back to the first-order update + :math:`\bar{\mathbf{D}}_n = \mathbf{D}_n`, which for :math:`\alpha_t = 1` is + algebraically identical to an explicit Euler step in :math:`\sigma`. + + The final step, where :math:`\sigma_{n-1} = 0`, likewise returns + :math:`\mathbf{D}_n` rather than the extrapolation + :math:`\bar{\mathbf{D}}_n`. This lowers the order of that one step relative + to a literal reading of Algorithm 2, and matches k-diffusion's terminal + handling and diffusers' ``lower_order_final`` behavior. + + The multistep extrapolation can amplify rounding error when two adjacent + timesteps are extremely close, since the coefficient then scales a difference + of successive data predictions that is itself near the precision of the + latent dtype. Very fine schedules may likewise become rounding-limited with + ``float16`` or ``bfloat16`` latents; prefer ``float32`` latent arithmetic in + that regime. + + .. warning:: + + The EDM parameterization is a **requirement**, not a default. Under any + other parameterization :math:`\mathbf{x} - t \cdot \text{RHS}` is not the + model's data prediction and :math:`-\ln t` is not the schedule's log-SNR + variable, so the update -- while still a consistent integrator -- is not + DPM-Solver++. :func:`~physicsnemo.diffusion.samplers.sample` rejects + incompatible noise schedulers; when calling :meth:`step` directly, it is + the caller's responsibility to supply an EDM schedule. + + .. note:: + + This solver is **stateful**: it caches the previous data prediction and + step size across calls to :meth:`step`. Call :meth:`reset` before + starting a new sampling trajectory when reusing an instance; + :func:`~physicsnemo.diffusion.samplers.sample` does this automatically. + A single instance cannot be shared by two interleaved trajectories. + + Parameters + ---------- + denoiser : Denoiser + A callable implementing the + :class:`~physicsnemo.diffusion.Denoiser` interface. Here it is + expected to return the right hand side of the ODE, + :math:`(\mathbf{x} - \mathbf{D}) / t`; the data prediction is recovered + internally as :math:`\mathbf{D} = \mathbf{x} - t \cdot \text{RHS}`. + Typically obtained via + :meth:`~physicsnemo.diffusion.noise_schedulers.NoiseScheduler.get_denoiser`, + but any callable with the correct signature can be used. + + Note + ---- + Reference: `DPM-Solver++: Fast Solver for Guided Sampling of Diffusion + Probabilistic Models `_, Algorithm 2. + + Examples + -------- + >>> import torch + >>> from physicsnemo.diffusion.samplers.solvers import DPMSolverPlusPlus2M + >>> + >>> denoiser = lambda x, t: x / (1 + t.view(-1, 1, 1, 1)**2) # Toy denoiser + >>> solver = DPMSolverPlusPlus2M(denoiser) + >>> x_t = torch.randn(1, 3, 8, 8) + >>> x_t = solver.step(x_t, torch.tensor([2.0]), torch.tensor([1.0])) + >>> x_t = solver.step(x_t, torch.tensor([1.0]), torch.tensor([0.5])) + >>> x_t.shape + torch.Size([1, 3, 8, 8]) + >>> isinstance(solver, Solver) + True + """ + + # Read by ``sample`` to reject incompatible noise schedulers; the class + # warning above explains why the EDM parameterization is required. + _requires_edm_parameterization = True + + # Opts this solver into the ``reset`` lifecycle call in ``sample``. Gated on + # a marker rather than on ``hasattr(solver, "reset")`` so that a pre-existing + # user-defined solver with an unrelated ``reset`` method is left untouched. + _requires_state_reset = True + + # Multistep history, cleared by ``reset``. + _D_prev: Tensor | None + _h_prev: Tensor | None + + def __init__(self, denoiser: Denoiser) -> None: + self.denoiser = denoiser + self.reset() + + def reset(self) -> None: + r""" + Clear the multistep history, starting a fresh trajectory. + + Drops the cached data prediction and step size from the previous + :meth:`step` call, so that the next call uses the first-order fallback, + and releases the cached latent-sized tensor. + :func:`~physicsnemo.diffusion.samplers.sample` calls this automatically + at the start of every trajectory. Call it explicitly when driving + :meth:`step` in a hand-written loop and reusing the instance for more + than one trajectory. + + Parameters + ---------- + None + + Returns + ------- + None + """ + self._D_prev = None + self._h_prev = None + + def step( + self, + x: Float[Tensor, " B *dims"], + t_cur: Float[Tensor, " B"], + t_next: Float[Tensor, " B"], + ) -> Float[Tensor, " B *dims"]: + r""" + Perform one DPM-Solver++(2M) integration step. + + Parameters + ---------- + x : Tensor + Current noisy latent state :math:`\mathbf{x}_{n}` of shape + :math:`(B, *)` where :math:`B` is the batch size. + t_cur : Tensor + Current diffusion time :math:`t_n` of shape :math:`(B,)`. + t_next : Tensor + Target diffusion time :math:`t_{n-1}` of shape :math:`(B,)`. + + Returns + ------- + Tensor + Updated latent state :math:`\mathbf{x}_{n-1}` at time + ``t_next``, same shape as ``x``. + """ + # Ensure contiguous strides so successive denoiser calls (across + # sampling steps) present the same stride layout to torch.compile, + # avoiding spurious recompilations / silently divergent traces. + t_cur = t_cur.contiguous() + t_next = t_next.contiguous() + + # Reshape t for broadcasting: (B,) -> (B, 1, ..., 1) + expected_shape = (-1,) + (1,) * (x.ndim - 1) + t_cur_bc = t_cur.reshape(expected_shape) + t_next_bc = t_next.reshape(expected_shape) + + # At t == 0 the state is already denoised and the denoiser is singular, + # since it returns (x - D) / t. Evaluate at a surrogate time to keep the + # call finite; the result is discarded, as such a step is the identity. + is_degenerate = t_cur_bc == 0 + t_cur_safe = torch.where(is_degenerate, torch.ones_like(t_cur_bc), t_cur_bc) + + # Single RHS evaluation; recover the data prediction D = x - t * RHS. + D = x - t_cur_safe * self.denoiser(x, t_cur_safe.reshape(t_cur.shape)) + + # Where t_next == 0 the extrapolation is dropped and the update returns D, + # matching the reference implementations. The surrogate keeps log() and the + # division by h finite on the unselected branch. + is_final = t_next_bc == 0 + t_next_safe = torch.where(is_final, 0.5 * t_cur_safe, t_next_bc) + + # Step size in lambda = -log(sigma), computed in at least float32: the + # coefficients below are conditioned on the ratio of successive sizes. + compute_dtype = torch.promote_types(t_cur_bc.dtype, torch.float32) + ratio_hp = t_next_safe.to(compute_dtype) / t_cur_safe.to(compute_dtype) + h = -torch.log(ratio_hp) + ratio = ratio_hp.to(x.dtype) + + if self._D_prev is None or self._h_prev is None: + D_bar = D # first-order fallback (exact exponential / DDIM) + else: + # Extrapolate the data prediction in lambda to lambda_cur + h/2: + # D + (1 / 2r) * (D - D_prev), r = h_prev / h, written h / (2 * h_prev). + # A repeated timestep gives h_prev == 0. Use a finite dummy + # denominator because torch.where evaluates both branches; the zero + # coefficient then falls back to first order. + has_history = self._h_prev != 0 + h_prev_safe = torch.where( + has_history, self._h_prev, torch.ones_like(self._h_prev) + ) + coeff = torch.where( + has_history, h / (2.0 * h_prev_safe), torch.zeros_like(h) + ).to(x.dtype) + D_bar = (1.0 + coeff) * D - coeff * self._D_prev + + x_next = torch.where( + is_degenerate, + x, + torch.where(is_final, D, ratio * x + (1.0 - ratio) * D_bar), + ) + + # Cache unconditionally to avoid a data-dependent branch (device sync, + # breaks fullgraph); `reset` clears it. Terminal and degenerate steps ran + # on a surrogate time, so their step size is not a real one and is stored + # as zero; a repeated timestep naturally yields h == 0 already. + self._D_prev = D + self._h_prev = torch.where(is_final | is_degenerate, torch.zeros_like(h), h) + + return x_next diff --git a/test/diffusion/data/test_solvers_dpmpp_2m_1d_step.pth b/test/diffusion/data/test_solvers_dpmpp_2m_1d_step.pth new file mode 100644 index 0000000000000000000000000000000000000000..16a664ad2f842e8f01a84c95fa5db3817313ae98 GIT binary patch literal 2243 zcmbVOZERCj7(QFKwshkoYKY(jN!qOow(ECSvFB_J)Q>x2DTvifZcBUJ8t>ZP-rG6B zu_Z&m`N7m|Zm0{I7!2Y-Wg#=;9>NbD5EoF>OdSb{@neYr_XFcsqUUU5w#BXsZ}R3Q z=brPN=f3ZAo^ySSq67f78r&3>;C>)+qQr`EzLiUgY$VZ=NU)X`)*N9)iA&T(LQ=?> zc$)WmYQYn?7@fuK0thI{?rm(0Ym*n~iKrit?g{O>#R^Q9i;-Tuh87*;b22 zDAwx(DzZW(Geb?BpqO2v@$&JIv~_k+Ba{S+$Q+YU8c>jDIFS;ROi&V3TNuC;GYV$K zenHKcm=#RXeqqUha5&EMTv&?6W1{AfWQ0p{k&V%?q!BdLexXbf%KgGp(dc7T>c1*e zO~(l&O<&&f8Gxfe`HWA+lPw{MrNt%NdOem~sG^xNI(2Qbi+(K4{ESRsT_Hwu6=g$w zv?#z2 z{(cqJ5zYJ@v_M~g166g6$?MTBia@GCmWMYtqm9db@VCYd`0_{H@;wK)psMTN$Rqdl z!HKrJ>fG^@sBh{!=obY%QMMbqk5=J*>BH!ONq{#myMoOwEplGo4p$p0aOqdw2-fV8 z+Z(*7{74yo_rU@9^Dhjvul3<4V-F$U`On;ccdW(RA3Gx-99v#jcXk#2wYS$Dj&+_s zKRgZL7drgow`bvrS8tOWQuTOL7li6-y>g)5ChJ>2h2X+81nb{~J?i7`tD~E+y`%^4 zd+!dsa(5p(hEBk9>1ufBk0rQns77v{co<%Nrvq0w2jz2)eykrF!J97Jj@Ed-lns14 z9#aOfwd*8$$-fkjA6s_j@Rd&>9)D3T-&2Mg2Zr!bWf!gwzK^E6JMhJq`%w=w>NZRb z!`Hr6;f4?P;zpMOInpb!&2R**whhSF2PfrIBd?%!mH|2V)(-UX<3`-S{&hGe>_R&} ze+~tzv)yT_HM?APv$@u8bJ(mVi^F7d zIZX~%t;6KBTU~Z58FW&XZy(tEgPxoYnm(q85}tSdOjE0*FS25B_dqUYrY#w=@DS(fUDypQS=iS_F%g3@ zJgR7j!IW4{P>Jz@l{^qa_MSsD!N({8ON|xlY>?WL2y_S{KHI%uERX`WXXc+XbI*Ui z^WA^G`M+N-v$LltM@Q-xD~L*lSE|)1>EzLmT!VdzjWtJLcrL}xx@YL~44job8`|LT zvW()(Y-M&bM;V7UIPzr`GJ%XEvmQB>Wn}2IT6G4k(dC-BtBoqPQLReUWYAm#Hzdj6 zOdDL13^PpOa+!l;xO&aiQ&wa)r;HYTScpaQ0&@}!~)rB;)fo1@O9r4s+C*~Exg0%)V5$JY8d8)5JGB@3p0_h-3jH8ibNvmA=WH+@(gu`xM^K#D&N2UgV)0;weMnFw6AX$2k0Ls@C3 z#g3&z314)t#9aaN&|t`KaIe2VTDomN*q8dDq?;b-_$S5WYNV0O`Dq4mliJ|yc`F#M z=Qc7dK8!YRJ&e~ky~(^};UeEPuH-u>HQDR;9X!q}K*mkmNq==S{=}P$mYr9lOG|f= z6@gbV|5O7?b}J+g2g*^(?0WOLN3P61uS3|}+km^XqezhY3z9W(7B>qma3QS{e!u2> z=s+K!>RTrgj(i$`DtaB^SZytK8S*5&qU|tn=4-hB!&%(%?bkTt^(sdchYv6 z1DURe@Rpu9!rS+EnAv@X*;#T97yj5{ZoTi0x>oW@$A`%zD|;7+2C^U}FN8$jJwayM z?;wp;5y(Ct5uUdMa^VuQ|KxILyxGd^xz$8ed-7qMW*98ApZH;a8QG~T$L67P=p#!H z?0681+jcHyT5<%$RqTr%A8N-J;=PD}PYmHGTj91tC`g>%LGM{S5a;=$=)HZvhJ&hD za=KoMTOyh;|G-Az?FfaWV-m!3QIf@aF1j{+3g3-KgzL45^j$axiAVlG$`_nq8qs}N zcz88V{U(P{m-`bMQ5e2_@6?Hhoks}ci;iEafaYk&JSeE>V9T8{y5kVanPO3OHiWcR~|fi`VWe~=aqx&nXky;+-UTrTO3|h?@Bu7|A_0qDaGp+R{~Ey zpLwn13Y_)nfndnQyY(5E>2x7+Umk_JfH`9RXYSplL zGC}vQUF3TQ$T*8Q%Tusy{R6l-I1`Dxe}@9^c_HTQGw_ecj-)B{ zGs4?*JYiY32>EC|U{1RkR=Sml9YZVe3QspYGqnvWiZ7Wxf|tOqf;q6rB?gMc-l#BL z#x%cM$L#KG!H>f3L5+kWWuI@s79W{-;m3Zc#={-@-WW2!Snykv{XQQxlTxr`%@^C6 zPmsA?FOZ9SoS^30i}=#MF8pf62GlO3NNVu`e5d3H(jF)Peybd7&lf?%@&TN1{v;{O zE*39IOfxsV9th7bMKD-CD6YO!FBW_yLg}su<@Xwyh}-S($FnDK-=BK1uxS~Y)$tCp z+~~sp%tYwAcqg9W69R>gRuRr!ca%!CV7BvZ^z_Oow7NJOlXt!aRM<3~@2LZYJX7l% zX~9ukmO=XmnXgRW|3#GFGz`K=ny*Q1)L0$p4JB*C3XMuC@t%}WWOS@BE;2qoIx;dY zS`;IS5(r}iG4b)Sg2>3&nCO_;=y*|FeEf(@xs^w5=3Cu?4Y|yQCH!>q8SCaj?7_J8 zY;({=%(1o{HV#i}b?T`xCrkwsp~o7oU>tP&^w7U>`HARbtu%NXdXlwhQxkt`;S)i} zsx)vM=#^C$q9S^6b`9$jzA;(HYG!8P!IwDUE`QODWmvK0Yo*>P& s`;QFR0hV1q@0t-gQq=I1QQUfO~{MwhdUS0QV$6&A-DRFVb<4)MQ%mO3Lcn=>DYt5DdcJc5>F(b_J;&SASPmVYDF0`&-d>2q&EaSIGM@8ub z+#}>4mJvniDs^c(Dff%p_OCB2w3y9wA#1Xf*~JTOMYN4BN;eg7O3tBXd#GgWa*M6BfHm-S=C0LhBJZ}K@jQ)#eckK%Kg-AX zG~`s6CmOdaYA7(9jAf;C8LQPiG%=GTB~JjJWquTJm-8_`fk-Oa1A%3$SlYr@T&szg z_};(eW;U^`ndUjvs4s``I^tt|4&hXk2M)Zd(a_diAe9WE?6+ScYvUSJ>4E?hbG(H-;x!i<8@&Q zS_O-9KE+4LEE4?FdswjfAoO-;;YqbA=+ODi&{M~ttd1mje?~2);`)y$|J1?G84W~r z?R&5uZiEum^XNcWJQh7Sh#e6MXXlSSIHRBe?LVng*KHAE%VP}G4w!(kRgmKJDY!k+ zjk>1&4Edk^qCU~*Pn?@}koqq})y+R7lQV|xu;c4*oLA%wY`u0r@eVtHvKQyW)-WSE zo7)7v?N5-Zg^7T>o`yWZC8+#8*p<=Ofjf7e!+s&1K*##P#aAE1vbv2RS(oNI>T7Wg zy=x?G({`bkLu90>cqQ6e^FL2lu>j2zuEPTzZUxwC3-fC!5=d zQE(ag2E~D|dpgWt)(yG)qhVukIeEe{=qhvMlB^~BP>rUwt@L~C88E~;Qnnwa3Aw=HFP?#WTzM&(0jp+>u&Xx zBnkfY;Tm+--K2`?IH|rN*i25xl!D*nO-P#kjjJcXirjgV(Q~U_Myz}qfTg*!~Vm?e9lRx`Z%4>;M>vn#5$~!P1Na=*X10xZ|)4{4?5 z0ILgl3pALDw3?~o5{geuR;0wIrY6S6rz9$qlnHW0k~}FjRi2uhlANfJb06XplH(J2 zmgVd6{G+%s@uAmw@r3V89wVJp#&3*TcQ=QO#T;pib4KCK=bUn4%rTBU7J8)N$wxt- z^9=njmmP~f(#j%7p=WYMo0#|=3mpqOQl;UeK)XDG4twHQ#F46t7=@_k#_+^wze#KU>Rvy~4& NU2tdoC_aBU_b=cIqFDd{ literal 0 HcmV?d00001 diff --git a/test/diffusion/test_samplers.py b/test/diffusion/test_samplers.py index fd6ae9f477..dcb3b7700e 100644 --- a/test/diffusion/test_samplers.py +++ b/test/diffusion/test_samplers.py @@ -26,11 +26,17 @@ ) from physicsnemo.diffusion.noise_schedulers import ( EDMNoiseScheduler, + IDDPMNoiseScheduler, + StudentTEDMNoiseScheduler, VENoiseScheduler, VPNoiseScheduler, ) +from physicsnemo.diffusion.noise_schedulers.domain_parallel import ( + DomainParallelNoiseScheduler, +) from physicsnemo.diffusion.samplers import sample from physicsnemo.diffusion.samplers.solvers import ( + DPMSolverPlusPlus2M, EulerSolver, HeunSolver, ) @@ -1002,3 +1008,290 @@ def test_backward_through_guided_sampling( for p in model.parameters() ) assert has_grad + + +# ============================================================================= +# DPM-Solver++(2M) Sampler Integration +# ============================================================================= + + +@pytest.mark.usefixtures("deterministic_settings") +class TestDPMSolverPlusPlus2MSampling: + """sample() integration for the stateful multistep solver.""" + + SHAPE = (BATCH, 3, 8, 6) + + def _components(self, device, sched_cls=EDMNoiseScheduler, num_steps=NUM_STEPS): + return _make_sampling_components( + sched_cls, + {}, + self.SHAPE, + Conv2dX0Predictor, + {"channels": 3}, + device, + num_steps=num_steps, + ) + + def test_entry_reset_discards_hand_driven_state(self, device): + """sample() must clear history it did not create. + + The exit reset alone is not enough: a caller may drive ``step`` by hand + -- a documented use -- and then pass the same instance to ``sample``. + Without the reset at entry, that stale data prediction is extrapolated + into the first step of the new trajectory. + """ + scheduler, _, denoiser, xN = self._components(device, num_steps=6) + solver = DPMSolverPlusPlus2M(denoiser) + + # Prime with a non-terminal step from an unrelated trajectory. + solver.step( + xN, + torch.full((BATCH,), 40.0, device=device), + torch.full((BATCH,), 12.0, device=device), + ) + assert solver._D_prev is not None + + primed = sample(denoiser, xN, scheduler, 6, solver=solver) + fresh = sample(denoiser, xN, scheduler, 6, solver=DPMSolverPlusPlus2M(denoiser)) + torch.testing.assert_close(primed, fresh, rtol=0, atol=0) + + # The cached latent-sized tensor must not outlive the trajectory either. + assert solver._D_prev is None + assert solver._h_prev is None + + def test_reset_hook_requires_the_marker(self, device): + """sample() must not call reset() on a solver that did not opt in. + + ``reset`` is a common method name, so a pre-existing user-defined solver + may already have one meaning something else entirely. The lifecycle call + is gated on ``_requires_state_reset`` so that such a solver keeps its + previous behavior. + """ + scheduler, _, denoiser, xN = self._components(device) + + class _UnrelatedReset: + """Solver whose reset() is its own business, not sampler lifecycle.""" + + def __init__(self, den): + self.denoiser, self.reset_calls = den, 0 + + def reset(self): + self.reset_calls += 1 + + def step(self, x, t_cur, t_next): + shape = (-1,) + (1,) * (x.ndim - 1) + return x + ( + t_next.reshape(shape) - t_cur.reshape(shape) + ) * self.denoiser(x, t_cur) + + solver = _UnrelatedReset(denoiser) + assert not hasattr(solver, "_requires_state_reset") + sample(denoiser, xN, scheduler, NUM_STEPS, solver=solver) + assert solver.reset_calls == 0 + + # The shipped solver does opt in, so it is still reset. + opted_in = DPMSolverPlusPlus2M(denoiser) + assert opted_in._requires_state_reset is True + sample(denoiser, xN, scheduler, NUM_STEPS, solver=opted_in) + assert opted_in._D_prev is None + + def test_history_released_after_error(self, device): + """Cached history must also be released when sampling raises.""" + scheduler, _, denoiser, xN = self._components(device) + + calls = [] + + def failing_denoiser(x, t): + calls.append(1) + if len(calls) > 1: + raise RuntimeError("boom") + return denoiser(x, t) + + solver = DPMSolverPlusPlus2M(failing_denoiser) + with pytest.raises(RuntimeError, match="boom"): + sample(failing_denoiser, xN, scheduler, 4, solver=solver) + assert solver._D_prev is None + assert solver._h_prev is None + + def test_multistep_path_is_exercised(self, device): + """Guard against a trajectory that never leaves the first-order path. + + With ``num_steps=2`` the first step has no history and the second lands + on ``t = 0``, where the update returns the data prediction and discards + the extrapolation -- so the result is bit-equal to Euler and proves + nothing about the multistep coefficients. At least three steps are + needed for the second-order path to affect the output. + """ + scheduler, _, denoiser, xN = self._components(device, num_steps=4) + + def relative_gap(num_steps): + dpm = sample(denoiser, xN, scheduler, num_steps, solver="dpmpp_2m") + # The registry key and a hand-built instance must agree exactly: + # sample() constructs the solver itself on the string path, and only + # the instance path additionally has state to reset. + by_instance = sample( + denoiser, xN, scheduler, num_steps, solver=DPMSolverPlusPlus2M(denoiser) + ) + torch.testing.assert_close(dpm, by_instance, rtol=0, atol=0) + euler = sample(denoiser, xN, scheduler, num_steps, solver="euler") + return float((dpm - euler).abs().max() / dpm.abs().max()) + + # Two steps: identical up to floating-point round-off. + assert relative_gap(2) < 1e-4 + # Three steps: the multistep update reaches the output. Measured gap is + # ~5e-1, i.e. four orders of magnitude above the two-step round-off. + assert relative_gap(3) > 1e-2 + + @pytest.mark.parametrize("as_instance", [False, True], ids=["by_name", "instance"]) + def test_rejects_non_edm_scheduler(self, device, as_instance): + """Both dispatch paths must reject an incompatible parameterization.""" + for sched_cls in (VENoiseScheduler, VPNoiseScheduler): + scheduler, _, denoiser, xN = self._components(device, sched_cls=sched_cls) + solver = DPMSolverPlusPlus2M(denoiser) if as_instance else "dpmpp_2m" + with pytest.raises(ValueError, match="EDM parameterization"): + sample(denoiser, xN, scheduler, NUM_STEPS, solver=solver) + + def test_accepts_wrapped_edm_scheduler(self, device): + """A scheduler wrapped for domain-parallel sampling is still EDM. + + The wrapper only changes tensor placement, so unwrapping it is required + or domain-parallel sampling would be rejected for no reason. + """ + scheduler, _, denoiser, xN = self._components(device) + + class _Wrapper: + """Stand-in for DomainParallelNoiseScheduler's public unwrap API. + + Deliberately has no ``__getattr__``: the real class delegates + explicitly rather than by fallback, so the capability is *not* + readable on the wrapper itself. Without the unwrap in + ``_check_edm_parameterization`` this scheduler would be rejected. + """ + + def __init__(self, inner): + self._inner = inner + + @property + def inner_scheduler(self): + return self._inner + + def timesteps(self, *args, **kwargs): + return self._inner.timesteps(*args, **kwargs) + + assert hasattr(DomainParallelNoiseScheduler, "inner_scheduler") + + wrapped = _Wrapper(scheduler) + assert not getattr(wrapped, "is_edm_parameterization", False) + + out = sample(denoiser, xN, wrapped, NUM_STEPS, solver="dpmpp_2m") + assert torch.isfinite(out).all() + + @pytest.mark.usefixtures("nop_compile") + def test_compiled_sample(self, device): + """The whole trajectory compiles fullgraph and reuses its graph.""" + torch._dynamo.config.error_on_recompile = False + torch._dynamo.reset() + + scheduler, _, denoiser, xN = self._components(device, num_steps=4) + solver = DPMSolverPlusPlus2M(denoiser) + + def do_sample(x): + return sample(denoiser, x, scheduler, 4, solver=solver) + + compiled = torch.compile(do_sample, fullgraph=True) + with torch.no_grad(): + first = compiled(xN) + torch._dynamo.config.error_on_recompile = True + try: + with torch.no_grad(): + second = compiled(xN) + finally: + torch._dynamo.config.error_on_recompile = False + + assert torch.isfinite(first).all() + torch.testing.assert_close(first, second, rtol=0, atol=0) + + # Compiling must not change the trajectory. Comparing the two compiled + # runs above only shows the graph is reused, not that it is right. + eager = sample(denoiser, xN, scheduler, 4, solver=DPMSolverPlusPlus2M(denoiser)) + torch.testing.assert_close(first, eager, rtol=1e-4, atol=1e-4) + + @pytest.mark.parametrize("guidance_config", GUIDANCE_CONFIGS) + def test_dps_guidance_reuse_is_clean(self, device, guidance_config): + """Guided sampling must be repeatable on one solver instance. + + DPS denoisers attach a per-step autograd graph, so the cached data + prediction is the one place a guided run could retain state or leak a + graph into the next trajectory. + """ + scheduler, _, denoiser, xN = _make_sampling_components( + EDMNoiseScheduler, + {}, + self.SHAPE, + Conv2dX0Predictor, + {"channels": 3}, + device, + num_steps=4, + guidance_config=guidance_config, + ) + solver = DPMSolverPlusPlus2M(denoiser) + + with torch.no_grad(): + first = sample(denoiser, xN, scheduler, 4, solver=solver) + second = sample(denoiser, xN, scheduler, 4, solver=solver) + + assert first.shape == self.SHAPE + assert torch.isfinite(first).all() + torch.testing.assert_close(first, second, rtol=0, atol=0) + assert solver._D_prev is None + assert solver._h_prev is None + # Under no_grad the result must not carry a graph from the guidance. + assert not first.requires_grad + assert first.grad_fn is None + + @pytest.mark.parametrize( + "sched_cls,sched_kwargs", + [(IDDPMNoiseScheduler, {}), (StudentTEDMNoiseScheduler, {})], + ids=["iddpm", "student_t_edm"], + ) + def test_accepts_non_inheriting_edm_scheduler( + self, device, sched_cls, sched_kwargs + ): + """Schedulers declaring the parameterization without inheriting it. + + IDDPM and Student-t EDM both satisfy sigma(t) = t and alpha(t) = 1 + without deriving from EDMNoiseScheduler -- they differ only in their + timestep ladder and latent distribution -- so an inheritance-based + check would reject them incorrectly. + """ + scheduler, _, denoiser, xN = _make_sampling_components( + sched_cls, + sched_kwargs, + self.SHAPE, + Conv2dX0Predictor, + {"channels": 3}, + device, + num_steps=4, + ) + out = sample(denoiser, xN, scheduler, 4, solver="dpmpp_2m") + assert out.shape == self.SHAPE + assert torch.isfinite(out).all() + + @pytest.mark.parametrize( + "sched_cls", [VENoiseScheduler, VPNoiseScheduler], ids=["ve", "vp"] + ) + @pytest.mark.usefixtures("nop_compile") + def test_rejects_non_edm_scheduler_when_compiled(self, device, sched_cls): + """Scheduler compatibility validation must survive fullgraph tracing.""" + torch._dynamo.config.error_on_recompile = False + torch._dynamo.reset() + + scheduler, _, denoiser, xN = self._components(device, sched_cls=sched_cls) + + def do_sample(x): + return sample(denoiser, x, scheduler, NUM_STEPS, solver="dpmpp_2m") + + with pytest.raises( + (ValueError, torch._dynamo.exc.Unsupported), match="EDM parameterization" + ): + torch.compile(do_sample, fullgraph=True)(xN) diff --git a/test/diffusion/test_solvers.py b/test/diffusion/test_solvers.py index 6265972c89..0a8e9385ca 100644 --- a/test/diffusion/test_solvers.py +++ b/test/diffusion/test_solvers.py @@ -16,11 +16,14 @@ """Tests for diffusion ODE/SDE solvers.""" +import math + import pytest import torch from physicsnemo.diffusion.noise_schedulers import EDMNoiseScheduler from physicsnemo.diffusion.samplers.solvers import ( + DPMSolverPlusPlus2M, EDMStochasticEulerSolver, EDMStochasticHeunSolver, EulerSolver, @@ -80,6 +83,10 @@ "stoch_heun_churn", True, ), + # Stateful: the parameterized golden files below take a single step from a + # fresh solver, so they pin only its first-order fallback. The multistep + # coefficients are covered by TestDPMSolverPlusPlus2M. + (DPMSolverPlusPlus2M, {}, "dpmpp_2m", False), ] @@ -329,6 +336,13 @@ def test_compiled_step( t_cur = torch.tensor([5.0] * shape[0], device=device) t_next = torch.tensor([2.5] * shape[0], device=device) + # Warm stateful solvers so this generic test covers steady-state graph + # reuse. Bootstrap compilation is covered by + # TestDPMSolverPlusPlus2M.test_compile_from_fresh_state_matches_eager. + if hasattr(solver, "reset"): + with torch.no_grad(): + solver.step(x, t_cur, t_next) + compiled_step = torch.compile(solver.step, fullgraph=True) with torch.no_grad(): @@ -347,3 +361,289 @@ def test_compiled_step( with torch.no_grad(): out_eager = solver.step(x, t_cur, t_next) torch.testing.assert_close(out_eager, out_compiled) + + +# ============================================================================= +# DPM-Solver++(2M) Specific Tests +# ============================================================================= + + +def _analytic_denoiser(scale: float = 1.0): + """ODE right-hand side for a Gaussian prior with standard deviation ``scale``. + + For :math:`p(x) = \\mathcal{N}(0, s^2)` the optimal denoiser is + :math:`D(x, t) = x s^2 / (s^2 + t^2)`, and the probability-flow ODE has the + closed-form solution :math:`x(t) = C \\sqrt{s^2 + t^2}`. This gives an exact + reference trajectory to measure the convergence order against. + """ + + def denoiser(x, t): + t_bc = t.reshape((-1,) + (1,) * (x.ndim - 1)) + D = x * scale**2 / (scale**2 + t_bc**2) + return (x - D) / t_bc + + return denoiser + + +def _exact_solution(x_init, t_init, t_final, scale=1.0): + """Exact PF-ODE solution for the Gaussian prior of ``_analytic_denoiser``.""" + return x_init * math.sqrt(scale**2 + t_final**2) / math.sqrt(scale**2 + t_init**2) + + +class TestDPMSolverPlusPlus2MConstructor: + """Tests for DPMSolverPlusPlus2M constructor.""" + + def test_attributes(self): + solver = DPMSolverPlusPlus2M(_identity_denoiser) + assert solver.denoiser is _identity_denoiser + assert isinstance(solver, Solver) + + +@pytest.mark.usefixtures("deterministic_settings") +class TestDPMSolverPlusPlus2M: + """Correctness, statefulness and compile behavior of DPM-Solver++(2M).""" + + def test_dbar_is_linear_extrapolation(self, device): + """The multistep coefficients must extrapolate D(lambda) to lambda + h/2. + + This pins the direction of the step-size ratio ``r = h_prev / h``. + The test is only sensitive to it on a *non-uniform* ladder: when + ``h_prev == h`` the correct and inverted coefficients coincide exactly. + """ + denoiser = _analytic_denoiser() + solver = DPMSolverPlusPlus2M(denoiser) + + # h_prev = log(2), h = log(4): deliberately non-uniform in lambda. + t0, t1, t2 = 8.0, 4.0, 1.0 + shape = (BATCH, 3, 8, 6) + x0 = make_input(shape, seed=7, device=device) + + def as_t(v): + return torch.full((BATCH,), v, device=device) + + def data_pred(x, t): + t_bc = torch.full((BATCH,) + (1,) * (len(shape) - 1), t, device=device) + return x - t_bc * denoiser(x, as_t(t)) + + D_prev = data_pred(x0, t0) + x1 = solver.step(x0, as_t(t0), as_t(t1)) + D = data_pred(x1, t1) + + # Linear interpolant through (lambda_prev, D_prev) and (lambda_cur, D), + # evaluated at the midpoint lambda_cur + h / 2. + lam = [-math.log(t) for t in (t0, t1, t2)] + h_prev, h = lam[1] - lam[0], lam[2] - lam[1] + coeff = h / (2.0 * h_prev) + D_bar = (1.0 + coeff) * D - coeff * D_prev + + ratio = t2 / t1 + expected = ratio * x1 + (1.0 - ratio) * D_bar + + x2 = solver.step(x1, as_t(t1), as_t(t2)) + torch.testing.assert_close(x2, expected, rtol=1e-5, atol=1e-6) + + def test_first_step_matches_euler(self, device): + """With no history the update reduces to an explicit Euler step in sigma.""" + denoiser = _analytic_denoiser() + shape = (BATCH, 3, 8, 6) + x = make_input(shape, seed=11, device=device) + t_cur = torch.full((BATCH,), 5.0, device=device) + t_next = torch.full((BATCH,), 2.5, device=device) + + dpm = DPMSolverPlusPlus2M(denoiser).step(x, t_cur, t_next) + euler = EulerSolver(denoiser).step(x, t_cur, t_next) + torch.testing.assert_close(dpm, euler, rtol=1e-5, atol=1e-6) + + @staticmethod + def _integrate(solver_cls, num_steps, device, scale=1.0, t_max=80.0, t_min=2e-3): + """Integrate the analytic PF-ODE on a Karras rho=7 ladder.""" + rho = 7.0 + ts = [ + ( + t_max ** (1 / rho) + + i / (num_steps - 1) * (t_min ** (1 / rho) - t_max ** (1 / rho)) + ) + ** rho + for i in range(num_steps) + ] + solver = solver_cls(_analytic_denoiser(scale)) + x = make_input((1, 4), seed=3, device=device) * t_max + for t_cur, t_next in zip(ts[:-1], ts[1:]): + x = solver.step( + x, + torch.full((1,), t_cur, device=device), + torch.full((1,), t_next, device=device), + ) + exact = _exact_solution( + make_input((1, 4), seed=3, device=device) * t_max, t_max, t_min, scale + ) + return float((x - exact).abs().max()) + + def test_convergence_and_accuracy_vs_euler(self, device): + """Verify asymptotic convergence and better accuracy than Euler at equal NFE. + + Only the asymptotic range is asserted: on a coarse ladder a wrong scheme + can be accidentally more accurate through error cancellation, so a + single-step-count threshold would not discriminate. + """ + errors = [ + self._integrate(DPMSolverPlusPlus2M, n, device) for n in (8, 16, 32, 64) + ] + assert all(errors[i + 1] < errors[i] for i in range(len(errors) - 1)), ( + f"error did not decrease monotonically: {errors}" + ) + + euler = self._integrate(EulerSolver, 64, device) + assert errors[-1] < euler / 5.0, ( + f"dpmpp_2m={errors[-1]:.3e} vs euler={euler:.3e}" + ) + + def test_reset_restores_first_order_path(self, device): + denoiser = _analytic_denoiser() + solver = DPMSolverPlusPlus2M(denoiser) + shape = (BATCH, 3, 8, 6) + x_init = make_input(shape, seed=5, device=device) + ts = [40.0, 12.0, 3.0, 0.4] + + def run(): + solver.reset() + x = x_init + for t_cur, t_next in zip(ts[:-1], ts[1:]): + x = solver.step( + x, + torch.full((BATCH,), t_cur, device=device), + torch.full((BATCH,), t_next, device=device), + ) + return x + + first, second = run(), run() + torch.testing.assert_close(first, second, rtol=0, atol=0) + + # Without the reset the stale history changes the result. + x = x_init + for t_cur, t_next in zip(ts[:-1], ts[1:]): + x = solver.step( + x, + torch.full((BATCH,), t_cur, device=device), + torch.full((BATCH,), t_next, device=device), + ) + assert not torch.allclose(x, first) + + def test_one_denoiser_call_per_step(self, device): + calls = [] + inner = _analytic_denoiser() + + def counting_denoiser(x, t): + calls.append(1) + return inner(x, t) + + solver = DPMSolverPlusPlus2M(counting_denoiser) + x = make_input((BATCH, 3, 8, 6), seed=13, device=device) + ts = [40.0, 12.0, 3.0, 0.4, 0.0] + for t_cur, t_next in zip(ts[:-1], ts[1:]): + x = solver.step( + x, + torch.full((BATCH,), t_cur, device=device), + torch.full((BATCH,), t_next, device=device), + ) + assert len(calls) == len(ts) - 1 + + def test_rounded_duplicate_timesteps_stay_finite(self, device): + """A ladder that collides only after casting must not poison the cache. + + Two adjacent timesteps that round to the same low-precision value give a + zero step size in lambda, which makes the multistep coefficient + singular. That step must be the identity and the next must fall back to + first order. This is how the case arises in practice: ``sample`` casts + the timesteps to the latent dtype, so a fine ladder loses distinctions + the schedule intended. Exact repeats also occur at full precision -- the + iDDPM ladder yields them at large step counts even in float64. + """ + denoiser = _analytic_denoiser() + solver = DPMSolverPlusPlus2M(denoiser) + x = make_input((BATCH, 3, 8, 6), seed=17, device=device).to(torch.bfloat16) + + # Strictly decreasing in float32; ts[1] and ts[2] collide in bfloat16. + ts = [40.0, 12.01, 12.0, 3.0, 0.4, 0.0] + cast = [ + torch.full((BATCH,), t, device=device, dtype=torch.bfloat16) for t in ts + ] + assert all(a > b for a, b in zip(ts[:-1], ts[1:])), "ladder must decrease" + assert torch.equal(cast[1], cast[2]), "ts[1] and ts[2] must collide in bfloat16" + + collision = 1 # index of the step whose endpoints collide after casting + for i, (t_cur, t_next) in enumerate(zip(cast[:-1], cast[1:])): + x_prev = x + x = solver.step(x, t_cur, t_next) + assert torch.isfinite(x).all(), ( + f"non-finite at step {i} ({ts[i]}->{ts[i + 1]})" + ) + if i == collision: + # A zero-length step must be exactly the identity. + torch.testing.assert_close(x, x_prev, rtol=0, atol=0) + + def test_zero_time_is_the_identity_and_differentiable(self, device): + """Repeated and zero timesteps stay finite and differentiable. + + The denoiser returns ``(x - D) / t`` and is singular at zero, so it runs + on a surrogate time and its result is discarded; the step is the + identity. A ladder rounded to a low-precision dtype can underflow to + zero before its final entry, so this is reachable in practice. + """ + denoiser = _analytic_denoiser() + solver = DPMSolverPlusPlus2M(denoiser) + x = make_input((BATCH, 3, 8, 6), seed=29, device=device).requires_grad_(True) + + out = x + for t_cur, t_next in zip( + [40.0, 12.0, 12.0, 3.0, 0.0], [12.0, 12.0, 3.0, 0.0, 0.0] + ): + before = out + out = solver.step( + out, + torch.full((BATCH,), t_cur, device=device), + torch.full((BATCH,), t_next, device=device), + ) + assert torch.isfinite(out).all(), f"non-finite at t={t_cur}->{t_next}" + + # The final entry steps from t_cur == 0, which must be an exact no-op. + torch.testing.assert_close(out, before, rtol=0, atol=0) + + out.sum().backward() + assert torch.isfinite(x.grad).all() + + @pytest.mark.usefixtures("nop_compile") + def test_compile_from_fresh_state_matches_eager(self, device): + """Compiling ``step`` from the reset state must reproduce an eager run. + + The shared TestStepCompile warms the solver first, so the first-order + bootstrap is only ever traced here. + """ + torch._dynamo.reset() + denoiser = _analytic_denoiser() + ts = [80.0, 30.0, 10.0, 3.0, 1.0, 0.3, 0.0] + + def run(solver, compile_step): + # step_fn must be bound to the same solver that is reset here, or + # the trajectory runs on another instance's history. + step_fn = ( + torch.compile(solver.step, fullgraph=True) + if compile_step + else solver.step + ) + solver.reset() + x = make_input((BATCH, 3, 8, 6), seed=23, device=device) + for t_cur, t_next in zip(ts[:-1], ts[1:]): + x = step_fn( + x, + torch.full((BATCH,), t_cur, device=device), + torch.full((BATCH,), t_next, device=device), + ) + return x + + with torch.no_grad(): + compiled = run(DPMSolverPlusPlus2M(denoiser), compile_step=True) + eager = run(DPMSolverPlusPlus2M(denoiser), compile_step=False) + + assert torch.isfinite(compiled).all() + torch.testing.assert_close(compiled, eager, rtol=1e-5, atol=1e-6) From db4586eda5719622e590d574cf595eae2088e222 Mon Sep 17 00:00:00 2001 From: gertl Date: Thu, 6 Aug 2026 10:54:48 +0200 Subject: [PATCH 2/2] Generalize DPM-Solver++(2M) schedule support Inject alpha, sigma and their derivatives for string-dispatched solvers, enabling supported linear-Gaussian schedules without an EDM-specific guard. Passed solver instances remain caller-configured. Use sigma-based endpoint handling, alpha-scaled terminal prediction and expm1 for stable small steps. Cover VP/VE convergence, nonstandard terminal endpoints, missing schedule functions, and fullgraph tracing. Co-Authored-By: Claude Opus 5 Signed-off-by: gertl --- CHANGELOG.md | 2 +- docs/api/diffusion/samplers.rst | 5 +- .../noise_schedulers/noise_schedulers.py | 24 -- physicsnemo/diffusion/samplers/samplers.py | 70 ++--- physicsnemo/diffusion/samplers/solvers.py | 243 ++++++++++++----- test/diffusion/test_samplers.py | 244 +++++++++++------- test/diffusion/test_solvers.py | 114 +++++++- 7 files changed, 468 insertions(+), 234 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index fb1671ac3d..353aca03fd 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -216,7 +216,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Adds `DPMSolverPlusPlus2M` (string key `"dpmpp_2m"`) to `physicsnemo.diffusion.samplers`, a second-order multistep DPM-Solver++ that reuses the previous step's data prediction and so needs only one denoiser - evaluation per step. Requires a noise scheduler in the EDM parameterization. + evaluation per step. Works with general linear-Gaussian noise schedulers. ### Changed diff --git a/docs/api/diffusion/samplers.rst b/docs/api/diffusion/samplers.rst index b130c3e248..5fdeccbd9f 100644 --- a/docs/api/diffusion/samplers.rst +++ b/docs/api/diffusion/samplers.rst @@ -408,9 +408,8 @@ There are two ways to use solvers: Second-order with configurable stochastic noise injection. - ``"dpmpp_2m"`` --- :class:`~physicsnemo.diffusion.samplers.solvers.DPMSolverPlusPlus2M`. - Second-order multistep, one denoiser evaluation per step. Requires a noise - scheduler in the EDM parameterization (:math:`\sigma(t) = t`, - :math:`\alpha(t) = 1`). + Second-order multistep, one denoiser evaluation per step. Configured from + the noise scheduler, so it applies to general linear-Gaussian schedules. **Custom solvers** can be defined by implementing the :class:`~physicsnemo.diffusion.samplers.solvers.Solver` protocol: any object diff --git a/physicsnemo/diffusion/noise_schedulers/noise_schedulers.py b/physicsnemo/diffusion/noise_schedulers/noise_schedulers.py index f817cb8ec5..311d163507 100644 --- a/physicsnemo/diffusion/noise_schedulers/noise_schedulers.py +++ b/physicsnemo/diffusion/noise_schedulers/noise_schedulers.py @@ -324,16 +324,6 @@ class LinearGaussianNoiseScheduler(ABC, NoiseScheduler): - :meth:`init_latents`: Initialize latent state (sampling) - :meth:`get_denoiser`: Get ODE/SDE RHS (sampling) - Attributes - ---------- - is_edm_parameterization : bool, default=False - Whether the schedule satisfies :math:`\sigma(t) = t` and - :math:`\alpha(t) = 1`, so that the diffusion time is the noise level. - Solvers derived only for this case, such as - :class:`~physicsnemo.diffusion.samplers.solvers.DPMSolverPlusPlus2M`, - are rejected by :func:`~physicsnemo.diffusion.samplers.sample` - otherwise. Subclasses satisfying both identities may set it to ``True``. - Examples -------- **Example 1:** A minimal EDM-like noise schedule. Only the abstract methods @@ -345,10 +335,6 @@ class LinearGaussianNoiseScheduler(ABC, NoiseScheduler): ... ) >>> >>> class SimpleEDMScheduler(LinearGaussianNoiseScheduler): - ... # sigma(t) = t and alpha(t) = 1 below, so opt in to the solvers - ... # that require the EDM parameterization. - ... is_edm_parameterization = True - ... ... def __init__(self, sigma_min=0.002, sigma_max=80.0, rho=7.0): ... self.sigma_min = sigma_min ... self.sigma_max = sigma_max @@ -400,10 +386,6 @@ class LinearGaussianNoiseScheduler(ABC, NoiseScheduler): """ - # Kept off the NoiseScheduler protocol, which is runtime_checkable: a new - # required member would break isinstance for duck-typed schedulers. - is_edm_parameterization: bool = False - @abstractmethod def sigma( self, @@ -1320,8 +1302,6 @@ class EDMNoiseScheduler(LinearGaussianNoiseScheduler): torch.Size([4, 3]) """ - is_edm_parameterization: bool = True - def __init__( self, sigma_min: float = 0.002, @@ -1830,8 +1810,6 @@ class IDDPMNoiseScheduler(LinearGaussianNoiseScheduler): torch.Size([4, 3, 8, 8]) """ - is_edm_parameterization: bool = True - def __init__( self, sigma_min: float = 0.002, @@ -2306,8 +2284,6 @@ class StudentTEDMNoiseScheduler(LinearGaussianNoiseScheduler): torch.Size([4, 3, 8, 8]) """ - is_edm_parameterization: bool = True - def __init__( self, sigma_min: float = 0.002, diff --git a/physicsnemo/diffusion/samplers/samplers.py b/physicsnemo/diffusion/samplers/samplers.py index c709218b53..445aa8587f 100644 --- a/physicsnemo/diffusion/samplers/samplers.py +++ b/physicsnemo/diffusion/samplers/samplers.py @@ -46,43 +46,47 @@ } -def _check_edm_parameterization( - noise_scheduler: NoiseScheduler, solver: Solver -) -> None: - r"""Raise if ``noise_scheduler`` is not in the EDM parameterization. +_SCHEDULE_FNS = ("alpha_fn", "sigma_fn", "alpha_dot_fn", "sigma_dot_fn") - Compatibility is read from the scheduler's ``is_edm_parameterization`` - attribute rather than probed numerically, because converting a tensor - comparison to a Python boolean would break - ``torch.compile(..., fullgraph=True)``. + +def _schedule_fns(noise_scheduler: NoiseScheduler) -> Dict[str, Any]: + r"""Extract the linear-Gaussian schedule functions required by a solver. + + Solvers written for general linear-Gaussian schedules + :math:`\mathbf{x}_t = \alpha_t \mathbf{x}_0 + \sigma_t \boldsymbol{\epsilon}` + need :math:`\alpha`, :math:`\sigma` and their time derivatives: the first + two to advance the state, the derivatives to recover the data prediction + from the ODE right-hand side that the denoiser returns. Parameters ---------- noise_scheduler : NoiseScheduler - The scheduler to validate. A + The scheduler from which to obtain the schedule functions. A :class:`~physicsnemo.diffusion.noise_schedulers.DomainParallelNoiseScheduler` - is unwrapped and its inner scheduler is validated instead. - solver : Solver - The solver requiring the parameterization; used in the error message. + is unwrapped first, since it delegates rather than exposing these + methods itself. Returns ------- - None + Dict[str, Any] + Keyword arguments for the solver constructor. Raises ------ ValueError - If the scheduler is not declared to be in the EDM parameterization. + If the scheduler does not provide callable ``alpha``, ``sigma``, + ``alpha_dot`` and ``sigma_dot`` methods. """ - # Unwrap the domain-parallel adapter: it only changes tensor placement. inner = getattr(noise_scheduler, "inner_scheduler", noise_scheduler) - - if not getattr(inner, "is_edm_parameterization", False): + fns = {key: getattr(inner, key[: -len("_fn")], None) for key in _SCHEDULE_FNS} + missing = sorted(k[: -len("_fn")] for k, v in fns.items() if not callable(v)) + if missing: raise ValueError( - f"{type(solver).__name__} requires a noise scheduler in the EDM " - f"parameterization (sigma(t) = t, alpha(t) = 1), but " - f"{type(inner).__name__} does not set is_edm_parameterization=True." + f"{type(inner).__name__} does not provide {', '.join(missing)}, which " + "this solver needs to advance a general linear-Gaussian schedule and " + "to recover the data prediction from the denoiser output." ) + return fns def _maybe_replicate_timesteps( @@ -260,8 +264,8 @@ def denoiser( :class:`~physicsnemo.diffusion.samplers.solvers.EDMStochasticHeunSolver`. * ``"dpmpp_2m"``: Second-order multistep DPM-Solver++, using a single - denoiser evaluation per step. Requires a noise scheduler in the EDM - parameterization. + denoiser evaluation per step. Configured from the noise scheduler, so + it applies to general linear-Gaussian schedules. See :class:`~physicsnemo.diffusion.samplers.solvers.DPMSolverPlusPlus2M`. time_steps : Tensor | None, default=None @@ -274,7 +278,9 @@ def denoiser( Additional options passed to the solver constructor. Only used when ``solver`` is a string; must be empty when ``solver`` is a :class:`Solver` instance. See individual solver classes for available - options. + options. Solvers configured from the noise scheduler (currently + ``"dpmpp_2m"``) ignore any schedule functions given here: the scheduler + also provides the time-steps, so it is authoritative. time_eval : List[int] | None, default=None Indices of time-steps at which to return intermediate samples. Must contain values in ``range(0, num_steps)`` (or ``range(0, @@ -368,11 +374,6 @@ def denoiser( >>> >>> # Define a minimal EDM-like scheduler from scratch >>> class MinimalScheduler: - ... # sigma=t and alpha=1 below, so opt in to the solvers that require - ... # the EDM parameterization. Duck-typed schedulers can set this too: - ... # it is read defensively, not required by the protocol. - ... is_edm_parameterization = True - ... ... def timesteps(self, num_steps, *, device=None, dtype=None): ... return torch.linspace(1.0, 0.0, num_steps + 1, ... device=device, dtype=dtype) @@ -415,7 +416,13 @@ def denoiser( f"Unknown solver '{solver}'. Available solvers: {available}." ) solver_cls = SOLVERS[solver] - solver_ = solver_cls(denoiser, **solver_options) + # Copy so the caller's dict is never mutated, and let the scheduler be + # authoritative: the time-steps come from it, so schedule functions + # describing anything else would silently disagree with them. + options = dict(solver_options) + if getattr(solver_cls, "_requires_schedule_fns", False): + options.update(_schedule_fns(noise_scheduler)) + solver_ = solver_cls(denoiser, **options) else: # Assume solver is a Solver-like object with a step method if solver_options: @@ -444,11 +451,6 @@ def denoiser( # callers that intentionally backprop through sample() are unaffected. outer_grad_enabled = torch.is_grad_enabled() - # Reject incompatible solver/scheduler parameterizations before evaluating - # the denoiser. - if getattr(solver_, "_requires_edm_parameterization", False): - _check_edm_parameterization(noise_scheduler, solver_) - # Stateful solvers opt into a reset hook, cleared here so a reused instance # cannot leak state from a previous trajectory. Gated on the marker rather # than on the presence of a ``reset`` method: ``reset`` is a common name, and diff --git a/physicsnemo/diffusion/samplers/solvers.py b/physicsnemo/diffusion/samplers/solvers.py index 25409c13cf..66df0f0aa5 100644 --- a/physicsnemo/diffusion/samplers/solvers.py +++ b/physicsnemo/diffusion/samplers/solvers.py @@ -800,21 +800,35 @@ class DPMSolverPlusPlus2M(Solver): Unlike :class:`HeunSolver`, which attains second order with *two* denoiser evaluations per step, this solver attains second order with a *single* - evaluation per step by reusing the previous step's data prediction - (a linear-multistep, Adams--Bashforth-style scheme applied to the - probability-flow ODE in exponential-integrator form). For a fixed budget of - network function evaluations it can therefore take twice as many steps as a - second-order single-step method. - - This implementation targets the variance-exploding (EDM) parameterization, - where the diffusion time equals the noise level (:math:`t = \sigma`) and - :math:`\alpha_t = 1`. Writing :math:`\lambda = -\ln \sigma`, with step size - :math:`h = \lambda_{n-1} - \lambda_n` so that - :math:`e^{-h} = \sigma_{n-1} / \sigma_n`, the update is + evaluation per step by reusing the previous step's data prediction: an + exponential-integrator multistep update of the probability-flow ODE. For a + fixed budget of denoiser evaluations it can therefore take twice as many + steps as a second-order single-step method. + + The scheme applies to linear-Gaussian diffusion schedules + :math:`\mathbf{x}_t = \alpha_t \mathbf{x}_0 + \sigma_t \boldsymbol{\epsilon}` + whose schedule functions return finite, strictly positive :math:`\alpha` and + :math:`\sigma` at every step before the zero-noise endpoint, and for which + :math:`\mathrm{d}\lambda / \mathrm{d}t \neq 0` there, where + :math:`\lambda = \ln(\alpha / \sigma)`. Positivity is a requirement on the + returned values, not on the schedule in exact arithmetic: a schedule whose + :math:`\alpha` underflows to zero in the dtype of the time-steps is outside + this range. The recovery below divides by + :math:`\alpha \dot{\sigma} - \sigma \dot{\alpha}`, which equals + :math:`-\alpha \sigma \, \mathrm{d}\lambda / \mathrm{d}t`, so a + stationary :math:`\lambda` is not admissible even if it is monotonic. + With :math:`\lambda` (half the log-SNR) and step size + :math:`h = \lambda_{n-1} - \lambda_n`, so that .. math:: - \mathbf{x}_{n-1} = e^{-h}\, \mathbf{x}_n - + \left(1 - e^{-h}\right) \bar{\mathbf{D}}_n , + e^{-h} = \frac{\alpha_n \sigma_{n-1}}{\alpha_{n-1} \sigma_n} , + + the update is + + .. math:: + \mathbf{x}_{n-1} = \frac{\sigma_{n-1}}{\sigma_n}\, \mathbf{x}_n + + \alpha_{n-1} \left(1 - e^{-h}\right) + \bar{\mathbf{D}}_n , where the extrapolated data prediction is @@ -824,33 +838,65 @@ class DPMSolverPlusPlus2M(Solver): \qquad r = \frac{h_{n+1}}{h_n} , and :math:`\mathbf{D}_n` is the denoised (data-space) prediction at step - :math:`n`. On the first step -- and whenever no history is available -- the - scheme falls back to the first-order update + :math:`n`. On the first step -- and whenever no usable history is available, + including after a zero-length step in :math:`\lambda` -- the scheme falls + back to the first-order update :math:`\bar{\mathbf{D}}_n = \mathbf{D}_n`, which for :math:`\alpha_t = 1` is algebraically identical to an explicit Euler step in :math:`\sigma`. - The final step, where :math:`\sigma_{n-1} = 0`, likewise returns + The final step, where :math:`\sigma_{n-1} = 0`, likewise uses :math:`\mathbf{D}_n` rather than the extrapolation - :math:`\bar{\mathbf{D}}_n`. This lowers the order of that one step relative + :math:`\bar{\mathbf{D}}_n`. It returns + :math:`\alpha_{n-1} \mathbf{D}_n`, which is :math:`\mathbf{D}_n` only + when :math:`\alpha = 1`. This lowers the order of that one step relative to a literal reading of Algorithm 2, and matches k-diffusion's terminal handling and diffusers' ``lower_order_final`` behavior. - The multistep extrapolation can amplify rounding error when two adjacent - timesteps are extremely close, since the coefficient then scales a difference - of successive data predictions that is itself near the precision of the - latent dtype. Very fine schedules may likewise become rounding-limited with - ``float16`` or ``bfloat16`` latents; prefer ``float32`` latent arithmetic in - that regime. + The multistep extrapolation amplifies rounding error when the previous step + is much shorter than the current one in :math:`\lambda`, since the + coefficient :math:`h_n / (2 h_{n+1})` is then large and scales the difference + of successive data predictions. Very fine or strongly non-uniform schedules + may become rounding-limited with ``float16`` or ``bfloat16`` latents; prefer + ``float32`` latent arithmetic in that regime. + + .. note:: - .. warning:: + When selected by name, + :func:`~physicsnemo.diffusion.samplers.sample` configures the solver from + its noise scheduler. For direct construction, including passing an + instance to :func:`~physicsnemo.diffusion.samplers.sample`, omitting the + four schedule functions selects the EDM defaults + (:math:`\alpha_t = 1`, :math:`\sigma_t = t`); provide all four for + another schedule. Successive manual :meth:`step` calls must + belong to one uninterrupted trajectory. - The EDM parameterization is a **requirement**, not a default. Under any - other parameterization :math:`\mathbf{x} - t \cdot \text{RHS}` is not the - model's data prediction and :math:`-\ln t` is not the schedule's log-SNR - variable, so the update -- while still a consistent integrator -- is not - DPM-Solver++. :func:`~physicsnemo.diffusion.samplers.sample` rejects - incompatible noise schedulers; when calling :meth:`step` directly, it is - the caller's responsibility to supply an EDM schedule. + The four schedule functions are supplied together or not at all; a partial + set is rejected, since it would combine a custom schedule with the EDM + defaults (:math:`\alpha_t = 1`, :math:`\sigma_t = t`) for the rest. The + signatures of the ``denoiser`` and of the four schedule functions are: + + .. code-block:: python + + def denoiser( + x: Tensor, # shape: (B, *dims) + t: Tensor, # shape: (B,) + ) -> Tensor: ... # ODE right-hand side, same shape as x + + def alpha_fn( + t: Tensor, # shape: (B, 1, ..., 1) + ) -> Tensor: ... # alpha_t, broadcastable to the shape of t + + def sigma_fn( + t: Tensor, # shape: (B, 1, ..., 1) + ) -> Tensor: ... # sigma_t, broadcastable to the shape of t + + def alpha_dot_fn( + t: Tensor, # shape: (B, 1, ..., 1) + ) -> Tensor: ... # d(alpha_t)/dt, broadcastable to the shape of t + + def sigma_dot_fn( + t: Tensor, # shape: (B, 1, ..., 1) + ) -> Tensor: ... # d(sigma_t)/dt, broadcastable to the shape of t .. note:: @@ -865,12 +911,26 @@ class DPMSolverPlusPlus2M(Solver): denoiser : Denoiser A callable implementing the :class:`~physicsnemo.diffusion.Denoiser` interface. Here it is - expected to return the right hand side of the ODE, - :math:`(\mathbf{x} - \mathbf{D}) / t`; the data prediction is recovered - internally as :math:`\mathbf{D} = \mathbf{x} - t \cdot \text{RHS}`. + expected to return the right-hand side of the probability-flow ODE. + The data prediction is recovered internally as + + .. math:: + \mathbf{D} = \frac{\dot{\sigma} \mathbf{x} - \sigma\, + \text{RHS}}{\alpha \dot{\sigma} - \sigma \dot{\alpha}} , + + which for the EDM schedule reduces to + :math:`\mathbf{x} - t \cdot \text{RHS}`. Typically obtained via :meth:`~physicsnemo.diffusion.noise_schedulers.NoiseScheduler.get_denoiser`, but any callable with the correct signature can be used. + alpha_fn : Callable[[Tensor], Tensor] | None, default=None + The schedule coefficient :math:`\alpha_t`. + sigma_fn : Callable[[Tensor], Tensor] | None, default=None + The noise level :math:`\sigma_t`. + alpha_dot_fn : Callable[[Tensor], Tensor] | None, default=None + The derivative :math:`\dot{\alpha}_t`. + sigma_dot_fn : Callable[[Tensor], Tensor] | None, default=None + The derivative :math:`\dot{\sigma}_t`. Note ---- @@ -893,21 +953,59 @@ class DPMSolverPlusPlus2M(Solver): True """ - # Read by ``sample`` to reject incompatible noise schedulers; the class - # warning above explains why the EDM parameterization is required. - _requires_edm_parameterization = True - # Opts this solver into the ``reset`` lifecycle call in ``sample``. Gated on # a marker rather than on ``hasattr(solver, "reset")`` so that a pre-existing # user-defined solver with an unrelated ``reset`` method is left untouched. _requires_state_reset = True + # Tells ``sample`` to inject the schedule functions of its noise scheduler. + # Kept off the ``Solver`` protocol, which is runtime checkable: a new member + # would make ``isinstance`` fail for solvers that implement only ``step``. + _requires_schedule_fns = True + # Multistep history, cleared by ``reset``. _D_prev: Tensor | None _h_prev: Tensor | None - def __init__(self, denoiser: Denoiser) -> None: + def __init__( + self, + denoiser: Denoiser, + *, + alpha_fn: Callable[[Tensor], Tensor] | None = None, + sigma_fn: Callable[[Tensor], Tensor] | None = None, + alpha_dot_fn: Callable[[Tensor], Tensor] | None = None, + sigma_dot_fn: Callable[[Tensor], Tensor] | None = None, + ) -> None: self.denoiser = denoiser + provided = [alpha_fn, sigma_fn, alpha_dot_fn, sigma_dot_fn] + if any(f is not None for f in provided) and not all( + f is not None for f in provided + ): + missing = [ + name + for name, f in zip( + ("alpha_fn", "sigma_fn", "alpha_dot_fn", "sigma_dot_fn"), provided + ) + if f is None + ] + raise ValueError( + "The schedule functions must be given together or not at all, " + f"but {', '.join(missing)} " + + ("is" if len(missing) == 1 else "are") + + " missing. Supplying only some of them would silently combine a " + "custom schedule with the EDM defaults for the rest." + ) + + # Default to the EDM schedule (alpha = 1, sigma = t), which makes the + # update reduce to the variance-exploding form. + self.alpha_fn = alpha_fn if alpha_fn is not None else torch.ones_like + self.sigma_fn = sigma_fn if sigma_fn is not None else (lambda t: t) + self.alpha_dot_fn = ( + alpha_dot_fn if alpha_dot_fn is not None else torch.zeros_like + ) + self.sigma_dot_fn = ( + sigma_dot_fn if sigma_dot_fn is not None else torch.ones_like + ) self.reset() def reset(self) -> None: @@ -969,30 +1067,46 @@ def step( t_cur_bc = t_cur.reshape(expected_shape) t_next_bc = t_next.reshape(expected_shape) - # At t == 0 the state is already denoised and the denoiser is singular, - # since it returns (x - D) / t. Evaluate at a surrogate time to keep the - # call finite; the result is discarded, as such a step is the identity. - is_degenerate = t_cur_bc == 0 - t_cur_safe = torch.where(is_degenerate, torch.ones_like(t_cur_bc), t_cur_bc) - - # Single RHS evaluation; recover the data prediction D = x - t * RHS. - D = x - t_cur_safe * self.denoiser(x, t_cur_safe.reshape(t_cur.shape)) + a_cur, s_cur = self.alpha_fn(t_cur_bc), self.sigma_fn(t_cur_bc) + a_next, s_next = self.alpha_fn(t_next_bc), self.sigma_fn(t_next_bc) + a_dot, s_dot = self.alpha_dot_fn(t_cur_bc), self.sigma_dot_fn(t_cur_bc) - # Where t_next == 0 the extrapolation is dropped and the update returns D, - # matching the reference implementations. The surrogate keeps log() and the - # division by h finite on the unselected branch. - is_final = t_next_bc == 0 - t_next_safe = torch.where(is_final, 0.5 * t_cur_safe, t_next_bc) + # At sigma == 0 the state is noise-free, and the right-hand side may be + # singular there. Evaluate at a surrogate time to keep the call finite; + # the result is dropped below, as such a step is the identity. + is_degenerate = s_cur == 0 + t_cur_safe = torch.where(is_degenerate, torch.ones_like(t_cur_bc), t_cur_bc) + a_cur = torch.where(is_degenerate, torch.ones_like(a_cur), a_cur) + s_cur = torch.where(is_degenerate, torch.ones_like(s_cur), s_cur) + a_dot = torch.where(is_degenerate, torch.zeros_like(a_dot), a_dot) + s_dot = torch.where(is_degenerate, torch.ones_like(s_dot), s_dot) + + # Single evaluation; recover the data prediction from the right-hand + # side. Written to avoid an explicit division by sigma near the + # endpoint. For the EDM schedule this is x - t * RHS. + rhs = self.denoiser(x, t_cur_safe.reshape(t_cur.shape)) + D = (s_dot * x - s_cur * rhs) / (a_cur * s_dot - s_cur * a_dot) + + # Step in lambda = log(alpha / sigma), half the log-SNR. lambda diverges + # at the zero-noise endpoint, so the general branch is evaluated with a + # strictly positive surrogate and the final mask selects alpha_next * D. + is_final = s_next == 0 + s_next_safe = torch.where(is_final, 0.5 * s_cur, s_next) - # Step size in lambda = -log(sigma), computed in at least float32: the - # coefficients below are conditioned on the ratio of successive sizes. compute_dtype = torch.promote_types(t_cur_bc.dtype, torch.float32) - ratio_hp = t_next_safe.to(compute_dtype) / t_cur_safe.to(compute_dtype) - h = -torch.log(ratio_hp) - ratio = ratio_hp.to(x.dtype) + emh_hp = (s_next_safe.to(compute_dtype) * a_cur.to(compute_dtype)) / ( + a_next.to(compute_dtype) * s_cur.to(compute_dtype) + ) + h = -torch.log(emh_hp) + # expm1 rather than 1 - exp(-h): the latter cancels catastrophically for + # the small h of a fine ladder, losing most of the value in bfloat16. + one_minus_emh = (-torch.expm1(-h)).to(x.dtype) + # Computed in the same promoted dtype. On non-terminal EDM steps it + # equals exp(-h); matching precision avoids rounding differences. + sigma_ratio = (s_next.to(compute_dtype) / s_cur.to(compute_dtype)).to(x.dtype) if self._D_prev is None or self._h_prev is None: - D_bar = D # first-order fallback (exact exponential / DDIM) + D_bar = D # first-order DPM-Solver++ update else: # Extrapolate the data prediction in lambda to lambda_cur + h/2: # D + (1 / 2r) * (D - D_prev), r = h_prev / h, written h / (2 * h_prev). @@ -1008,16 +1122,17 @@ def step( ).to(x.dtype) D_bar = (1.0 + coeff) * D - coeff * self._D_prev + general = sigma_ratio * x + a_next * one_minus_emh * D_bar + # At sigma_next == 0, use the lower-order final update alpha_next * D + # instead of the extrapolated D_bar; see the class docstring. x_next = torch.where( - is_degenerate, - x, - torch.where(is_final, D, ratio * x + (1.0 - ratio) * D_bar), + is_degenerate, x, torch.where(is_final, a_next * D, general) ) # Cache unconditionally to avoid a data-dependent branch (device sync, - # breaks fullgraph); `reset` clears it. Terminal and degenerate steps ran - # on a surrogate time, so their step size is not a real one and is stored - # as zero; a repeated timestep naturally yields h == 0 already. + # breaks fullgraph); `reset` clears it. Terminal and degenerate steps use + # a surrogate step size, which is not a real one, so store zero; a + # repeated timestep naturally yields h == 0 already. self._D_prev = D self._h_prev = torch.where(is_final | is_degenerate, torch.zeros_like(h), h) diff --git a/test/diffusion/test_samplers.py b/test/diffusion/test_samplers.py index dcb3b7700e..2dd1d738de 100644 --- a/test/diffusion/test_samplers.py +++ b/test/diffusion/test_samplers.py @@ -26,14 +26,9 @@ ) from physicsnemo.diffusion.noise_schedulers import ( EDMNoiseScheduler, - IDDPMNoiseScheduler, - StudentTEDMNoiseScheduler, VENoiseScheduler, VPNoiseScheduler, ) -from physicsnemo.diffusion.noise_schedulers.domain_parallel import ( - DomainParallelNoiseScheduler, -) from physicsnemo.diffusion.samplers import sample from physicsnemo.diffusion.samplers.solvers import ( DPMSolverPlusPlus2M, @@ -1091,7 +1086,6 @@ def step(self, x, t_cur, t_next): # The shipped solver does opt in, so it is still reset. opted_in = DPMSolverPlusPlus2M(denoiser) - assert opted_in._requires_state_reset is True sample(denoiser, xN, scheduler, NUM_STEPS, solver=opted_in) assert opted_in._D_prev is None @@ -1118,7 +1112,8 @@ def test_multistep_path_is_exercised(self, device): With ``num_steps=2`` the first step has no history and the second lands on ``t = 0``, where the update returns the data prediction and discards - the extrapolation -- so the result is bit-equal to Euler and proves + the extrapolation -- so the result matches Euler up to floating-point + round-off and proves nothing about the multistep coefficients. At least three steps are needed for the second-order path to affect the output. """ @@ -1138,54 +1133,9 @@ def relative_gap(num_steps): # Two steps: identical up to floating-point round-off. assert relative_gap(2) < 1e-4 - # Three steps: the multistep update reaches the output. Measured gap is - # ~5e-1, i.e. four orders of magnitude above the two-step round-off. + # Three steps: the multistep update reaches the output. assert relative_gap(3) > 1e-2 - @pytest.mark.parametrize("as_instance", [False, True], ids=["by_name", "instance"]) - def test_rejects_non_edm_scheduler(self, device, as_instance): - """Both dispatch paths must reject an incompatible parameterization.""" - for sched_cls in (VENoiseScheduler, VPNoiseScheduler): - scheduler, _, denoiser, xN = self._components(device, sched_cls=sched_cls) - solver = DPMSolverPlusPlus2M(denoiser) if as_instance else "dpmpp_2m" - with pytest.raises(ValueError, match="EDM parameterization"): - sample(denoiser, xN, scheduler, NUM_STEPS, solver=solver) - - def test_accepts_wrapped_edm_scheduler(self, device): - """A scheduler wrapped for domain-parallel sampling is still EDM. - - The wrapper only changes tensor placement, so unwrapping it is required - or domain-parallel sampling would be rejected for no reason. - """ - scheduler, _, denoiser, xN = self._components(device) - - class _Wrapper: - """Stand-in for DomainParallelNoiseScheduler's public unwrap API. - - Deliberately has no ``__getattr__``: the real class delegates - explicitly rather than by fallback, so the capability is *not* - readable on the wrapper itself. Without the unwrap in - ``_check_edm_parameterization`` this scheduler would be rejected. - """ - - def __init__(self, inner): - self._inner = inner - - @property - def inner_scheduler(self): - return self._inner - - def timesteps(self, *args, **kwargs): - return self._inner.timesteps(*args, **kwargs) - - assert hasattr(DomainParallelNoiseScheduler, "inner_scheduler") - - wrapped = _Wrapper(scheduler) - assert not getattr(wrapped, "is_edm_parameterization", False) - - out = sample(denoiser, xN, wrapped, NUM_STEPS, solver="dpmpp_2m") - assert torch.isfinite(out).all() - @pytest.mark.usefixtures("nop_compile") def test_compiled_sample(self, device): """The whole trajectory compiles fullgraph and reuses its graph.""" @@ -1216,6 +1166,147 @@ def do_sample(x): eager = sample(denoiser, xN, scheduler, 4, solver=DPMSolverPlusPlus2M(denoiser)) torch.testing.assert_close(first, eager, rtol=1e-4, atol=1e-4) + @pytest.mark.parametrize( + "sched_cls", + [VPNoiseScheduler, VENoiseScheduler], + ids=["vp", "ve"], + ) + def test_string_dispatch_configures_linear_gaussian_schedule( + self, device, sched_cls + ): + """Selecting by name must configure the solver from the scheduler. + + Shape and finiteness alone would not detect the schedule functions being + dropped, since the EDM defaults also produce a finite result. The string + result is therefore compared against an explicitly configured instance, + and shown to differ from the EDM defaults. Only non-EDM schedules are + used: on EDM the injected functions equal the defaults, so the + comparison would hold even if the injection were removed. + """ + scheduler, _, denoiser, xN = self._components( + device, sched_cls=sched_cls, num_steps=6 + ) + by_name = sample(denoiser, xN, scheduler, 6, solver="dpmpp_2m") + configured = sample( + denoiser, + xN, + scheduler, + 6, + solver=DPMSolverPlusPlus2M( + denoiser, + alpha_fn=scheduler.alpha, + sigma_fn=scheduler.sigma, + alpha_dot_fn=scheduler.alpha_dot, + sigma_dot_fn=scheduler.sigma_dot, + ), + ) + torch.testing.assert_close(by_name, configured, rtol=0, atol=0) + + # The defaults integrate a different ODE here, so the equality above + # would be vacuous if the schedule functions were ignored. + edm_defaults = sample( + denoiser, xN, scheduler, 6, solver=DPMSolverPlusPlus2M(denoiser) + ) + assert not torch.allclose(by_name, edm_defaults) + + def test_schedule_functions_taken_from_wrapped_scheduler(self, device): + """A wrapper delegating to an inner scheduler must be unwrapped. + + ``DomainParallelNoiseScheduler`` exposes the schedule functions only via + its inner scheduler, so the wrapped result must equal the unwrapped one. + """ + scheduler, _, denoiser, xN = self._components(device) + + class _Wrapper: + """Minimal wrapper exposing inner_scheduler and timesteps.""" + + def __init__(self, inner): + self._inner = inner + + @property + def inner_scheduler(self): + return self._inner + + def timesteps(self, *args, **kwargs): + return self._inner.timesteps(*args, **kwargs) + + wrapped = sample( + denoiser, xN, _Wrapper(scheduler), NUM_STEPS, solver="dpmpp_2m" + ) + direct = sample(denoiser, xN, scheduler, NUM_STEPS, solver="dpmpp_2m") + torch.testing.assert_close(wrapped, direct, rtol=0, atol=0) + + def test_scheduler_without_schedule_functions_is_rejected(self, device): + """A scheduler missing the schedule functions must fail clearly.""" + scheduler, _, denoiser, xN = self._components(device) + + class _Incomplete: + def __init__(self, inner): + self._inner = inner + + def timesteps(self, *args, **kwargs): + return self._inner.timesteps(*args, **kwargs) + + def alpha(self, t): + return torch.ones_like(t) + + with pytest.raises(ValueError, match="sigma"): + sample(denoiser, xN, _Incomplete(scheduler), NUM_STEPS, solver="dpmpp_2m") + + class _NoScheduleFns(_Incomplete): + alpha = None + + # Providing none of the four must be rejected here rather than falling + # through to the EDM defaults, which would silently integrate the wrong + # ODE. The constructor's partial-set guard cannot catch this case. + with pytest.raises(ValueError, match="alpha, alpha_dot, sigma, sigma_dot"): + sample( + denoiser, xN, _NoScheduleFns(scheduler), NUM_STEPS, solver="dpmpp_2m" + ) + + def test_scheduler_overrides_conflicting_solver_options(self, device): + """The scheduler is authoritative and the caller's dict is untouched. + + The time-steps come from the scheduler, so schedule functions describing + anything else would silently disagree with them. + """ + scheduler, _, denoiser, xN = self._components(device) + conflicting = { + "alpha_fn": lambda t: torch.full_like(t, 2.0), + "sigma_fn": lambda t: 3.0 * t, + "alpha_dot_fn": torch.zeros_like, + "sigma_dot_fn": lambda t: torch.full_like(t, 3.0), + } + opts = dict(conflicting) + out = sample( + denoiser, xN, scheduler, NUM_STEPS, solver="dpmpp_2m", solver_options=opts + ) + assert opts.keys() == conflicting.keys() + assert all(opts[k] is conflicting[k] for k in conflicting) + + expected = sample(denoiser, xN, scheduler, NUM_STEPS, solver="dpmpp_2m") + torch.testing.assert_close(out, expected, rtol=0, atol=0) + + @pytest.mark.parametrize( + "sched_cls", [VPNoiseScheduler, VENoiseScheduler], ids=["vp", "ve"] + ) + @pytest.mark.usefixtures("nop_compile") + def test_compiled_sample_with_general_schedule(self, device, sched_cls): + """Bound scheduler methods must remain traceable under fullgraph.""" + torch._dynamo.reset() + scheduler, _, denoiser, xN = self._components( + device, sched_cls=sched_cls, num_steps=4 + ) + + def do_sample(x): + return sample(denoiser, x, scheduler, 4, solver="dpmpp_2m") + + with torch.no_grad(): + compiled = torch.compile(do_sample, fullgraph=True)(xN) + eager = do_sample(xN) + assert torch.isfinite(compiled).all() + torch.testing.assert_close(compiled, eager, rtol=1e-4, atol=1e-4) + @pytest.mark.parametrize("guidance_config", GUIDANCE_CONFIGS) def test_dps_guidance_reuse_is_clean(self, device, guidance_config): """Guided sampling must be repeatable on one solver instance. @@ -1248,50 +1339,3 @@ def test_dps_guidance_reuse_is_clean(self, device, guidance_config): # Under no_grad the result must not carry a graph from the guidance. assert not first.requires_grad assert first.grad_fn is None - - @pytest.mark.parametrize( - "sched_cls,sched_kwargs", - [(IDDPMNoiseScheduler, {}), (StudentTEDMNoiseScheduler, {})], - ids=["iddpm", "student_t_edm"], - ) - def test_accepts_non_inheriting_edm_scheduler( - self, device, sched_cls, sched_kwargs - ): - """Schedulers declaring the parameterization without inheriting it. - - IDDPM and Student-t EDM both satisfy sigma(t) = t and alpha(t) = 1 - without deriving from EDMNoiseScheduler -- they differ only in their - timestep ladder and latent distribution -- so an inheritance-based - check would reject them incorrectly. - """ - scheduler, _, denoiser, xN = _make_sampling_components( - sched_cls, - sched_kwargs, - self.SHAPE, - Conv2dX0Predictor, - {"channels": 3}, - device, - num_steps=4, - ) - out = sample(denoiser, xN, scheduler, 4, solver="dpmpp_2m") - assert out.shape == self.SHAPE - assert torch.isfinite(out).all() - - @pytest.mark.parametrize( - "sched_cls", [VENoiseScheduler, VPNoiseScheduler], ids=["ve", "vp"] - ) - @pytest.mark.usefixtures("nop_compile") - def test_rejects_non_edm_scheduler_when_compiled(self, device, sched_cls): - """Scheduler compatibility validation must survive fullgraph tracing.""" - torch._dynamo.config.error_on_recompile = False - torch._dynamo.reset() - - scheduler, _, denoiser, xN = self._components(device, sched_cls=sched_cls) - - def do_sample(x): - return sample(denoiser, x, scheduler, NUM_STEPS, solver="dpmpp_2m") - - with pytest.raises( - (ValueError, torch._dynamo.exc.Unsupported), match="EDM parameterization" - ): - torch.compile(do_sample, fullgraph=True)(xN) diff --git a/test/diffusion/test_solvers.py b/test/diffusion/test_solvers.py index 0a8e9385ca..212328f96c 100644 --- a/test/diffusion/test_solvers.py +++ b/test/diffusion/test_solvers.py @@ -21,7 +21,11 @@ import pytest import torch -from physicsnemo.diffusion.noise_schedulers import EDMNoiseScheduler +from physicsnemo.diffusion.noise_schedulers import ( + EDMNoiseScheduler, + VENoiseScheduler, + VPNoiseScheduler, +) from physicsnemo.diffusion.samplers.solvers import ( DPMSolverPlusPlus2M, EDMStochasticEulerSolver, @@ -339,7 +343,7 @@ def test_compiled_step( # Warm stateful solvers so this generic test covers steady-state graph # reuse. Bootstrap compilation is covered by # TestDPMSolverPlusPlus2M.test_compile_from_fresh_state_matches_eager. - if hasattr(solver, "reset"): + if getattr(solver, "_requires_state_reset", False): with torch.no_grad(): solver.step(x, t_cur, t_next) @@ -369,11 +373,11 @@ def test_compiled_step( def _analytic_denoiser(scale: float = 1.0): - """ODE right-hand side for a Gaussian prior with standard deviation ``scale``. + r"""ODE right-hand side for a Gaussian data distribution with standard deviation ``scale``. - For :math:`p(x) = \\mathcal{N}(0, s^2)` the optimal denoiser is + For :math:`p(x) = \mathcal{N}(0, s^2)` the optimal denoiser is :math:`D(x, t) = x s^2 / (s^2 + t^2)`, and the probability-flow ODE has the - closed-form solution :math:`x(t) = C \\sqrt{s^2 + t^2}`. This gives an exact + closed-form solution :math:`x(t) = C \sqrt{s^2 + t^2}`. This gives an exact reference trajectory to measure the convergence order against. """ @@ -386,17 +390,23 @@ def denoiser(x, t): def _exact_solution(x_init, t_init, t_final, scale=1.0): - """Exact PF-ODE solution for the Gaussian prior of ``_analytic_denoiser``.""" + """Exact PF-ODE solution for the Gaussian data distribution of ``_analytic_denoiser``.""" return x_init * math.sqrt(scale**2 + t_final**2) / math.sqrt(scale**2 + t_init**2) class TestDPMSolverPlusPlus2MConstructor: """Tests for DPMSolverPlusPlus2M constructor.""" - def test_attributes(self): + def test_default_attributes(self): solver = DPMSolverPlusPlus2M(_identity_denoiser) assert solver.denoiser is _identity_denoiser assert isinstance(solver, Solver) + # Without schedule functions the solver uses the EDM schedule. + t = torch.tensor(3.0) + assert solver.alpha_fn(t) == torch.ones_like(t) + assert solver.sigma_fn(t) == t + assert solver.alpha_dot_fn(t) == torch.zeros_like(t) + assert solver.sigma_dot_fn(t) == torch.ones_like(t) @pytest.mark.usefixtures("deterministic_settings") @@ -582,7 +592,95 @@ def test_rounded_duplicate_timesteps_stay_finite(self, device): # A zero-length step must be exactly the identity. torch.testing.assert_close(x, x_prev, rtol=0, atol=0) - def test_zero_time_is_the_identity_and_differentiable(self, device): + @pytest.mark.parametrize( + "sched_cls", [VPNoiseScheduler, VENoiseScheduler], ids=["vp", "ve"] + ) + def test_converges_on_vp_and_ve_schedules(self, device, sched_cls): + """Second-order behavior is not specific to the EDM parameterization. + + Results are compared with the exact solution because solvers can use + different terminal updates. The window is deliberately narrow and the + assertions broad: coarse ladders are pre-asymptotic, and on VE the error + changes sign near 150 steps, so an order estimated across that crossing + is meaningless. + """ + sched = sched_cls() + scale = 1.0 + + def x0_predictor(x, t): + t_bc = t.reshape((-1,) + (1,) * (x.ndim - 1)) + a, sg = sched.alpha(t_bc), sched.sigma(t_bc) + return x * a * scale**2 / (a**2 * scale**2 + sg**2) + + denoiser = sched.get_denoiser(x0_predictor=x0_predictor, denoising_type="ode") + + errors = [] + for num_steps in (24, 32, 48, 64): + ts = sched.timesteps(num_steps, device=device, dtype=torch.float64) + a0, s0 = sched.alpha(ts[0]), sched.sigma(ts[0]) + aT, sT = sched.alpha(ts[-1]), sched.sigma(ts[-1]) + scale0 = float(torch.sqrt(a0**2 * scale**2 + s0**2)) + xT = make_input((1, 64), seed=41, device=device).double() * scale0 + exact = xT * float(torch.sqrt(aT**2 * scale**2 + sT**2)) / scale0 + + solver = DPMSolverPlusPlus2M( + denoiser, + alpha_fn=sched.alpha, + sigma_fn=sched.sigma, + alpha_dot_fn=sched.alpha_dot, + sigma_dot_fn=sched.sigma_dot, + ) + x = xT + for t_cur, t_next in zip(ts[:-1], ts[1:]): + x = solver.step(x, t_cur.expand(1), t_next.expand(1)) + errors.append(float((x - exact).abs().max() / exact.abs().max())) + + assert all(errors[i + 1] < errors[i] for i in range(len(errors) - 1)), ( + f"error did not decrease monotonically: {errors}" + ) + # Refining 24 -> 64 steps is a factor 8/3; a second-order method gains + # roughly (8/3)^2 ~ 7x. Assert well below that to stay robust. + assert errors[0] / errors[-1] > 3.0, f"convergence too slow: {errors}" + + def test_rejects_partial_schedule_functions(self): + """Partially supplied schedule functions must be rejected. + + Accepting a subset would silently combine a custom schedule with the + EDM defaults for the rest, which integrates a different ODE than the + caller intended. + """ + denoiser = _analytic_denoiser() + with pytest.raises(ValueError, match="together or not at all"): + DPMSolverPlusPlus2M(denoiser, sigma_fn=lambda t: t) + + def test_terminal_step_scales_the_data_prediction_by_alpha(self): + """The final step must return ``alpha_next * D``, not ``D``. + + Every shipped scheduler has ``alpha == 1`` at the zero-noise endpoint + and reaches ``sigma == 0`` only at ``t == 0``, so neither the factor + nor the detection on ``sigma`` rather than ``t`` is observable there. + This schedule has ``alpha(1) = 1.5`` and ``sigma(1) = 0``, separating + both. + """ + rhs = torch.full((1, 4), 0.25) + solver = DPMSolverPlusPlus2M( + lambda x, t: rhs, + alpha_fn=lambda t: 1.0 + t / 2.0, + sigma_fn=lambda t: t - 1.0, + alpha_dot_fn=lambda t: torch.full_like(t, 0.5), + sigma_dot_fn=torch.ones_like, + ) + x = torch.linspace(-1.0, 1.0, 4).reshape(1, 4) + # One ordinary step first: without history the extrapolation equals D, + # and the general branch would land on the same value as the terminal + # one, so detecting the final step on sigma would not be observable. + x = solver.step(x, torch.tensor([3.0]), torch.tensor([2.0])) + # D = (sigma_dot x - sigma rhs) / (alpha sigma_dot - sigma alpha_dot) + # is (x - rhs) / 1.5 at t = 2, so alpha_next * D is exactly x - rhs. + x_next = solver.step(x, torch.tensor([2.0]), torch.tensor([1.0])) + torch.testing.assert_close(x_next, x - rhs) + + def test_zero_sigma_is_the_identity_and_differentiable(self, device): """Repeated and zero timesteps stay finite and differentiable. The denoiser returns ``(x - D) / t`` and is singular at zero, so it runs