Skip to content
Closed
Show file tree
Hide file tree
Changes from 1 commit
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. Requires a noise scheduler in the EDM parameterization.

### Changed

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

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

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

Guidance
~~~~~~~~

Expand Down
24 changes: 24 additions & 0 deletions physicsnemo/diffusion/noise_schedulers/noise_schedulers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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``.

@CharlelieLrt CharlelieLrt Aug 5, 2026 •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why is this needed? The new solver should work for any noise schedule that is of the linear-gaussian form x_t = alpha_t * x_0 + sigma_t * epsilon. I think it would be better (and not very difficult) to generalize its implementation to make it work for any schedule (alpha_t, sigma_t) (and not just (1, t)), rather than adding edge cases and checks.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

You are right, the implementation can be generalized using the scheduler’s alpha, sigma and their derivatives alpha_dot and sigma_dot. These provide the schedule values needed by the update and allow the solver to recover the data prediction from the denoiser’s ODE output. I can let sample() pass them when constructing dpmpp_2m, making the string API work with every linear-Gaussian scheduler. I would then remove the EDM-only checks and detect the final step using sigma_next == 0.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, I think that would make sense.
For broader context, we have other noise schedules that we might merge soon, so we don't want to have this new solver compatible with EDM but not with those.

Just a detail:

I can let sample() pass them when constructing dpmpp_2m, making the string API work with every linear-Gaussian scheduler.

This would work, but is a little bit cumbersome in terms of API for sample(), because one would have to pass all sigma, alpha and their derivative to the sample() (which itself needs to pass them to the constructor of the solver). But that's more a detail, you can go for this and I'll try to think of a more elegant way to handle this in the sample() API.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The new commit generalizes the solver using alpha, sigma and their derivatives from the noise scheduler. When dpmpp_2m is selected by name, sample() configures it automatically. I have removed the EDM-only attribute and checks.

Examples
--------
**Example 1:** A minimal EDM-like noise schedule. Only the abstract methods
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -1302,6 +1320,8 @@ class EDMNoiseScheduler(LinearGaussianNoiseScheduler):
torch.Size([4, 3])
"""

is_edm_parameterization: bool = True

def __init__(
self,
sigma_min: float = 0.002,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
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
123 changes: 99 additions & 24 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,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"],
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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()
Loading
Loading