diff --git a/CHANGELOG.md b/CHANGELOG.md index 6d796875a5..353aca03fd 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. Works with general linear-Gaussian noise schedulers. ### Changed diff --git a/docs/api/diffusion/samplers.rst b/docs/api/diffusion/samplers.rst index 5b06ea5c58..5fdeccbd9f 100644 --- a/docs/api/diffusion/samplers.rst +++ b/docs/api/diffusion/samplers.rst @@ -406,6 +406,10 @@ 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. 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 @@ -502,6 +506,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/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..445aa8587f 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,53 @@ "heun": HeunSolver, "edm_stochastic_euler": EDMStochasticEulerSolver, "edm_stochastic_heun": EDMStochasticHeunSolver, + "dpmpp_2m": DPMSolverPlusPlus2M, } +_SCHEDULE_FNS = ("alpha_fn", "sigma_fn", "alpha_dot_fn", "sigma_dot_fn") + + +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 from which to obtain the schedule functions. A + :class:`~physicsnemo.diffusion.noise_schedulers.DomainParallelNoiseScheduler` + is unwrapped first, since it delegates rather than exposing these + methods itself. + + Returns + ------- + Dict[str, Any] + Keyword arguments for the solver constructor. + + Raises + ------ + ValueError + If the scheduler does not provide callable ``alpha``, ``sigma``, + ``alpha_dot`` and ``sigma_dot`` methods. + """ + inner = getattr(noise_scheduler, "inner_scheduler", noise_scheduler) + 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(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( t_steps: Float[Tensor, " N_plus_1"], xN: Float[Tensor, " B *dims"], @@ -78,7 +123,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 +263,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. 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 Optional 1D tensor of shape :math:`(N + 1,)` containing explicit diffusion time values :math:`t_N, t_{N-1}, ..., t_0` in decreasing @@ -226,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, @@ -362,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: @@ -391,6 +451,18 @@ def denoiser( # callers that intentionally backprop through sample() are unaffected. outer_grad_enabled = torch.is_grad_enabled() + # 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 +476,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..66df0f0aa5 100644 --- a/physicsnemo/diffusion/samplers/solvers.py +++ b/physicsnemo/diffusion/samplers/solvers.py @@ -792,3 +792,348 @@ 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: 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:: + 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 + + .. 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 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 uses + :math:`\mathbf{D}_n` rather than the extrapolation + :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 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:: + + 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 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:: + + 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 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 + ---- + 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 + """ + + # 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, + *, + 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: + 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) + + 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) + + # 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) + + compute_dtype = torch.promote_types(t_cur_bc.dtype, torch.float32) + 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 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). + # 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 + + 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, a_next * D, general) + ) + + # Cache unconditionally to avoid a data-dependent branch (device sync, + # 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) + + 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 0000000000..16a664ad2f Binary files /dev/null and b/test/diffusion/data/test_solvers_dpmpp_2m_1d_step.pth differ diff --git a/test/diffusion/data/test_solvers_dpmpp_2m_2d_step.pth b/test/diffusion/data/test_solvers_dpmpp_2m_2d_step.pth new file mode 100644 index 0000000000..2a2965d2f2 Binary files /dev/null and b/test/diffusion/data/test_solvers_dpmpp_2m_2d_step.pth differ diff --git a/test/diffusion/data/test_solvers_dpmpp_2m_3d_step.pth b/test/diffusion/data/test_solvers_dpmpp_2m_3d_step.pth new file mode 100644 index 0000000000..fdb86edbbb Binary files /dev/null and b/test/diffusion/data/test_solvers_dpmpp_2m_3d_step.pth differ diff --git a/test/diffusion/test_samplers.py b/test/diffusion/test_samplers.py index fd6ae9f477..2dd1d738de 100644 --- a/test/diffusion/test_samplers.py +++ b/test/diffusion/test_samplers.py @@ -31,6 +31,7 @@ ) from physicsnemo.diffusion.samplers import sample from physicsnemo.diffusion.samplers.solvers import ( + DPMSolverPlusPlus2M, EulerSolver, HeunSolver, ) @@ -1002,3 +1003,339 @@ 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) + 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 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. + """ + 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. + assert relative_gap(3) > 1e-2 + + @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( + "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. + + 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 diff --git a/test/diffusion/test_solvers.py b/test/diffusion/test_solvers.py index 6265972c89..212328f96c 100644 --- a/test/diffusion/test_solvers.py +++ b/test/diffusion/test_solvers.py @@ -16,11 +16,18 @@ """Tests for diffusion ODE/SDE solvers.""" +import math + 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, EDMStochasticHeunSolver, EulerSolver, @@ -80,6 +87,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 +340,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 getattr(solver, "_requires_state_reset", False): + 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 +365,383 @@ 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): + 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 + :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 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_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") +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) + + @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 + 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)