Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
12 changes: 12 additions & 0 deletions docs/api/diffusion/samplers.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -502,6 +506,14 @@ Solvers
:members:
:exclude-members: __init__

:code:`DPMSolverPlusPlus2M`
^^^^^^^^^^^^^^^^^^^^^^^^^^^

.. autoclass:: physicsnemo.diffusion.samplers.solvers.DPMSolverPlusPlus2M
:show-inheritance:
:members:
:exclude-members: __init__

Guidance
~~~~~~~~

Expand Down
1 change: 1 addition & 0 deletions physicsnemo/diffusion/samplers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
129 changes: 103 additions & 26 deletions physicsnemo/diffusion/samplers/samplers.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from physicsnemo.domain_parallel.shard_tensor import scatter_tensor

from .solvers import (
DPMSolverPlusPlus2M,
EDMStochasticEulerSolver,
EDMStochasticHeunSolver,
EulerSolver,
Expand All @@ -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"],
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand All @@ -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()
Loading
Loading