Skip to content

Commit c7e5ea0

Browse files
Fix diffusion inferers to support diffusers-style schedulers (#8870)
Fixes #8597. ### Description This PR adds compatibility for `diffusers`-style schedulers in MONAI's diffusion inferers (`DiffusionInferer` and `LatentDiffusionInferer`). Specifically, this includes: - **Scheduler Step Normalization**: The sampling path in `inferer.py` now detects `return_dict` support on `scheduler.step`, requests tuple output when available, and normalizes both MONAI tuple returns and `diffusers`-style `prev_sample` outputs. This shared helper also benefits the ControlNet diffusion inferers. - **Likelihood Path Fixes**: Fixed the likelihood paths to consistently use the specific `scheduler` instance passed into the call. - **Regression Testing**: Added regression coverage in `test_diffusion_inferer.py` and `test_latent_diffusion_inferer.py` using a small diffusers-style DDPM shim that returns a `prev_sample` output object by default. Verified locally (`31 passed` and `46 passed`). ### Types of changes <!--- Put an `x` in all the boxes that apply, and remove the not applicable items --> - [x] Non-breaking change (fix or new feature that would not break existing functionality). - [ ] Breaking change (fix or new feature that would cause existing functionality to change). - [x] New tests added to cover the changes. - [ ] Integration tests passed locally by running `./runtests.sh -f -u --net --coverage`. - [ ] Quick tests passed locally by running `./runtests.sh --quick --unittests --disttests`. - [ ] In-line docstrings updated. - [ ] Documentation updated, tested `make html` command in the `docs/` folder. --------- Signed-off-by: ugbotueferhire <ugbotueferhire@gmail.com> Signed-off-by: Eric Kerfoot <17726042+ericspod@users.noreply.github.com> Co-authored-by: Eric Kerfoot <17726042+ericspod@users.noreply.github.com>
1 parent 880a429 commit c7e5ea0

3 files changed

Lines changed: 248 additions & 36 deletions

File tree

‎monai/inferers/inferer.py‎

Lines changed: 149 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111

1212
from __future__ import annotations
1313

14+
import inspect
1415
import math
1516
import warnings
1617
from abc import ABC, abstractmethod
@@ -862,6 +863,96 @@ def __init__(self, scheduler: Scheduler) -> None: # type: ignore[override]
862863

863864
self.scheduler = scheduler
864865

866+
@staticmethod
867+
def _scheduler_step_supports_kwarg(scheduler: Scheduler, kwarg: str) -> bool:
868+
try:
869+
return kwarg in inspect.signature(scheduler.step).parameters
870+
except (TypeError, ValueError):
871+
return False
872+
873+
@staticmethod
874+
def _get_previous_sample_from_step_output(step_output: Any) -> torch.Tensor:
875+
if isinstance(step_output, tuple):
876+
return step_output[0]
877+
if isinstance(step_output, Mapping):
878+
return step_output["prev_sample"]
879+
if hasattr(step_output, "prev_sample"):
880+
return step_output.prev_sample
881+
raise TypeError("Unsupported scheduler.step output. Expected a tuple or an object with `prev_sample`.")
882+
883+
@staticmethod
884+
def _get_scheduler_name(scheduler: Scheduler) -> str:
885+
if hasattr(scheduler, "_get_name"):
886+
return scheduler._get_name()
887+
return scheduler.__class__.__name__
888+
889+
@staticmethod
890+
def _get_scheduler_config_value(scheduler: Scheduler, name: str, default: Any = None) -> Any:
891+
config = getattr(scheduler, "config", None)
892+
if isinstance(config, Mapping):
893+
if name in config:
894+
return config[name]
895+
elif config is not None and hasattr(config, name):
896+
return getattr(config, name)
897+
898+
if hasattr(scheduler, name):
899+
return getattr(scheduler, name)
900+
return default
901+
902+
@staticmethod
903+
def _get_posterior_mean(
904+
scheduler: Scheduler, timestep: int | torch.Tensor, x_0: torch.Tensor, x_t: torch.Tensor
905+
) -> torch.Tensor:
906+
alpha_t = scheduler.alphas[timestep]
907+
alpha_prod_t = scheduler.alphas_cumprod[timestep]
908+
alpha_prod_t_prev = scheduler.alphas_cumprod[timestep - 1] if timestep > 0 else scheduler.one
909+
910+
x_0_coefficient = alpha_prod_t_prev.sqrt() * scheduler.betas[timestep] / (1 - alpha_prod_t)
911+
x_t_coefficient = alpha_t.sqrt() * (1 - alpha_prod_t_prev) / (1 - alpha_prod_t)
912+
913+
return x_0_coefficient * x_0 + x_t_coefficient * x_t
914+
915+
def _get_posterior_variance(
916+
self, scheduler: Scheduler, timestep: int | torch.Tensor, predicted_variance: torch.Tensor | None = None
917+
) -> torch.Tensor:
918+
alpha_prod_t = scheduler.alphas_cumprod[timestep]
919+
alpha_prod_t_prev = scheduler.alphas_cumprod[timestep - 1] if timestep > 0 else scheduler.one
920+
variance = (1 - alpha_prod_t_prev) / (1 - alpha_prod_t) * scheduler.betas[timestep]
921+
variance_type = self._get_scheduler_config_value(scheduler, "variance_type")
922+
923+
if variance_type == "fixed_small":
924+
variance = torch.clamp(variance, min=1e-20)
925+
elif variance_type == "fixed_large":
926+
variance = scheduler.betas[timestep]
927+
elif variance_type == "learned" and predicted_variance is not None:
928+
return predicted_variance
929+
elif variance_type == "learned_range" and predicted_variance is not None:
930+
min_log = variance
931+
max_log = scheduler.betas[timestep]
932+
frac = (predicted_variance + 1) / 2
933+
variance = frac * max_log + (1 - frac) * min_log
934+
935+
return variance
936+
937+
def _scheduler_step(
938+
self,
939+
scheduler: Scheduler,
940+
model_output: torch.Tensor,
941+
timestep: int | torch.Tensor,
942+
sample: torch.Tensor,
943+
next_timestep: int | torch.Tensor | None = None,
944+
) -> torch.Tensor:
945+
step_kwargs = {}
946+
if self._scheduler_step_supports_kwarg(scheduler, "return_dict"):
947+
step_kwargs["return_dict"] = False
948+
949+
if isinstance(scheduler, RFlowScheduler):
950+
step_output = scheduler.step(model_output, timestep, sample, next_timestep, **step_kwargs) # type: ignore
951+
else:
952+
step_output = scheduler.step(model_output, timestep, sample, **step_kwargs) # type: ignore
953+
954+
return self._get_previous_sample_from_step_output(step_output)
955+
865956
def __call__( # type: ignore[override]
866957
self,
867958
inputs: torch.Tensor,
@@ -941,7 +1032,12 @@ def sample(
9411032
scheduler = self.scheduler
9421033
image = input_noise
9431034

944-
all_next_timesteps = torch.cat((scheduler.timesteps[1:], torch.tensor([0], dtype=scheduler.timesteps.dtype)))
1035+
all_next_timesteps = torch.cat(
1036+
(
1037+
scheduler.timesteps[1:],
1038+
torch.tensor([0], dtype=scheduler.timesteps.dtype, device=scheduler.timesteps.device),
1039+
)
1040+
)
9451041
if verbose and has_tqdm:
9461042
progress_bar = tqdm(
9471043
zip(scheduler.timesteps, all_next_timesteps),
@@ -985,10 +1081,9 @@ def sample(
9851081
model_output = model_output_uncond + cfg * (model_output_cond - model_output_uncond)
9861082

9871083
# 2. compute previous image: x_t -> x_t-1
988-
if not isinstance(scheduler, RFlowScheduler):
989-
image, _ = scheduler.step(model_output, t, image) # type: ignore
990-
else:
991-
image, _ = scheduler.step(model_output, t, image, next_t) # type: ignore
1084+
image = self._scheduler_step(
1085+
scheduler=scheduler, model_output=model_output, timestep=t, sample=image, next_timestep=next_t
1086+
)
9921087
if save_intermediates and t % intermediate_steps == 0:
9931088
intermediates.append(image)
9941089

@@ -1029,10 +1124,10 @@ def get_likelihood(
10291124

10301125
if not scheduler:
10311126
scheduler = self.scheduler
1032-
if scheduler._get_name() != "DDPMScheduler":
1127+
scheduler_name = self._get_scheduler_name(scheduler)
1128+
if scheduler_name != "DDPMScheduler":
10331129
raise NotImplementedError(
1034-
f"Likelihood computation is only compatible with DDPMScheduler,"
1035-
f" you are using {scheduler._get_name()}"
1130+
f"Likelihood computation is only compatible with DDPMScheduler, you are using {scheduler_name},"
10361131
)
10371132
if mode not in ["crossattn", "concat"]:
10381133
raise NotImplementedError(f"{mode} condition is not supported")
@@ -1047,7 +1142,7 @@ def get_likelihood(
10471142
total_kl = torch.zeros(inputs.shape[0]).to(inputs.device)
10481143
for t in progress_bar:
10491144
timesteps = torch.full(inputs.shape[:1], t, device=inputs.device).long()
1050-
noisy_image = self.scheduler.add_noise(original_samples=inputs, noise=noise, timesteps=timesteps)
1145+
noisy_image = scheduler.add_noise(original_samples=inputs, noise=noise, timesteps=timesteps)
10511146
diffusion_model = (
10521147
partial(diffusion_model, seg=seg)
10531148
if isinstance(diffusion_model, SPADEDiffusionModelUNet)
@@ -1060,7 +1155,8 @@ def get_likelihood(
10601155
model_output = diffusion_model(x=noisy_image, timesteps=timesteps, context=conditioning)
10611156

10621157
# get the model's predicted mean, and variance if it is predicted
1063-
if model_output.shape[1] == inputs.shape[1] * 2 and scheduler.variance_type in ["learned", "learned_range"]:
1158+
variance_type = self._get_scheduler_config_value(scheduler, "variance_type")
1159+
if model_output.shape[1] == inputs.shape[1] * 2 and variance_type in ["learned", "learned_range"]:
10641160
model_output, predicted_variance = torch.split(model_output, inputs.shape[1], dim=1)
10651161
else:
10661162
predicted_variance = None
@@ -1073,15 +1169,17 @@ def get_likelihood(
10731169

10741170
# 2. compute predicted original sample from predicted noise also called
10751171
# "predicted x_0" of formula (15) from https://arxiv.org/pdf/2006.11239.pdf
1076-
if scheduler.prediction_type == "epsilon":
1172+
prediction_type = self._get_scheduler_config_value(scheduler, "prediction_type")
1173+
if prediction_type == "epsilon":
10771174
pred_original_sample = (noisy_image - beta_prod_t ** (0.5) * model_output) / alpha_prod_t ** (0.5)
1078-
elif scheduler.prediction_type == "sample":
1175+
elif prediction_type == "sample":
10791176
pred_original_sample = model_output
1080-
elif scheduler.prediction_type == "v_prediction":
1177+
elif prediction_type == "v_prediction":
10811178
pred_original_sample = (alpha_prod_t**0.5) * noisy_image - (beta_prod_t**0.5) * model_output
10821179
# 3. Clip "predicted x_0"
1083-
if scheduler.clip_sample:
1084-
pred_original_sample = torch.clamp(pred_original_sample, -1, 1)
1180+
if self._get_scheduler_config_value(scheduler, "clip_sample"):
1181+
clip_sample_range = self._get_scheduler_config_value(scheduler, "clip_sample_range", 1.0)
1182+
pred_original_sample = torch.clamp(pred_original_sample, -clip_sample_range, clip_sample_range)
10851183

10861184
# 4. Compute coefficients for pred_original_sample x_0 and current sample x_t
10871185
# See formula (7) from https://arxiv.org/pdf/2006.11239.pdf
@@ -1093,11 +1191,15 @@ def get_likelihood(
10931191
predicted_mean = pred_original_sample_coeff * pred_original_sample + current_sample_coeff * noisy_image
10941192

10951193
# get the posterior mean and variance
1096-
posterior_mean = scheduler._get_mean(timestep=t, x_0=inputs, x_t=noisy_image) # type: ignore[operator]
1097-
posterior_variance = scheduler._get_variance(timestep=t, predicted_variance=predicted_variance) # type: ignore[operator]
1194+
posterior_mean = self._get_posterior_mean(scheduler=scheduler, timestep=t, x_0=inputs, x_t=noisy_image)
1195+
posterior_variance = self._get_posterior_variance(
1196+
scheduler=scheduler, timestep=t, predicted_variance=predicted_variance
1197+
)
10981198

10991199
log_posterior_variance = torch.log(posterior_variance)
1100-
log_predicted_variance = torch.log(predicted_variance) if predicted_variance else log_posterior_variance
1200+
log_predicted_variance = (
1201+
torch.log(predicted_variance) if predicted_variance is not None else log_posterior_variance
1202+
)
11011203

11021204
if t == 0:
11031205
# compute -log p(x_0|x_1)
@@ -1510,7 +1612,12 @@ def sample( # type: ignore[override]
15101612
scheduler = self.scheduler
15111613
image = input_noise
15121614

1513-
all_next_timesteps = torch.cat((scheduler.timesteps[1:], torch.tensor([0], dtype=scheduler.timesteps.dtype)))
1615+
all_next_timesteps = torch.cat(
1616+
(
1617+
scheduler.timesteps[1:],
1618+
torch.tensor([0], dtype=scheduler.timesteps.dtype, device=scheduler.timesteps.device),
1619+
)
1620+
)
15141621
if verbose and has_tqdm:
15151622
progress_bar = tqdm(
15161623
zip(scheduler.timesteps, all_next_timesteps),
@@ -1584,10 +1691,9 @@ def sample( # type: ignore[override]
15841691
model_output = model_output_uncond + cfg * (model_output_cond - model_output_uncond)
15851692

15861693
# 3. compute previous image: x_t -> x_t-1
1587-
if not isinstance(scheduler, RFlowScheduler):
1588-
image, _ = scheduler.step(model_output, t, image) # type: ignore
1589-
else:
1590-
image, _ = scheduler.step(model_output, t, image, next_t) # type: ignore
1694+
image = self._scheduler_step(
1695+
scheduler=scheduler, model_output=model_output, timestep=t, sample=image, next_timestep=next_t
1696+
)
15911697

15921698
if save_intermediates and t % intermediate_steps == 0:
15931699
intermediates.append(image)
@@ -1632,10 +1738,10 @@ def get_likelihood( # type: ignore[override]
16321738

16331739
if not scheduler:
16341740
scheduler = self.scheduler
1635-
if scheduler._get_name() != "DDPMScheduler":
1741+
scheduler_name = self._get_scheduler_name(scheduler)
1742+
if scheduler_name != "DDPMScheduler":
16361743
raise NotImplementedError(
1637-
f"Likelihood computation is only compatible with DDPMScheduler,"
1638-
f" you are using {scheduler._get_name()}"
1744+
f"Likelihood computation is only compatible with DDPMScheduler, you are using {scheduler_name}."
16391745
)
16401746
if mode not in ["crossattn", "concat"]:
16411747
raise NotImplementedError(f"{mode} condition is not supported")
@@ -1648,7 +1754,7 @@ def get_likelihood( # type: ignore[override]
16481754
total_kl = torch.zeros(inputs.shape[0]).to(inputs.device)
16491755
for t in progress_bar:
16501756
timesteps = torch.full(inputs.shape[:1], t, device=inputs.device).long()
1651-
noisy_image = self.scheduler.add_noise(original_samples=inputs, noise=noise, timesteps=timesteps)
1757+
noisy_image = scheduler.add_noise(original_samples=inputs, noise=noise, timesteps=timesteps)
16521758

16531759
diffuse = diffusion_model
16541760
if isinstance(diffusion_model, SPADEDiffusionModelUNet):
@@ -1681,7 +1787,8 @@ def get_likelihood( # type: ignore[override]
16811787
mid_block_additional_residual=mid_block_res_sample,
16821788
)
16831789
# get the model's predicted mean, and variance if it is predicted
1684-
if model_output.shape[1] == inputs.shape[1] * 2 and scheduler.variance_type in ["learned", "learned_range"]:
1790+
variance_type = self._get_scheduler_config_value(scheduler, "variance_type")
1791+
if model_output.shape[1] == inputs.shape[1] * 2 and variance_type in ["learned", "learned_range"]:
16851792
model_output, predicted_variance = torch.split(model_output, inputs.shape[1], dim=1)
16861793
else:
16871794
predicted_variance = None
@@ -1694,15 +1801,17 @@ def get_likelihood( # type: ignore[override]
16941801

16951802
# 2. compute predicted original sample from predicted noise also called
16961803
# "predicted x_0" of formula (15) from https://arxiv.org/pdf/2006.11239.pdf
1697-
if scheduler.prediction_type == "epsilon":
1804+
prediction_type = self._get_scheduler_config_value(scheduler, "prediction_type")
1805+
if prediction_type == "epsilon":
16981806
pred_original_sample = (noisy_image - beta_prod_t ** (0.5) * model_output) / alpha_prod_t ** (0.5)
1699-
elif scheduler.prediction_type == "sample":
1807+
elif prediction_type == "sample":
17001808
pred_original_sample = model_output
1701-
elif scheduler.prediction_type == "v_prediction":
1809+
elif prediction_type == "v_prediction":
17021810
pred_original_sample = (alpha_prod_t**0.5) * noisy_image - (beta_prod_t**0.5) * model_output
17031811
# 3. Clip "predicted x_0"
1704-
if scheduler.clip_sample:
1705-
pred_original_sample = torch.clamp(pred_original_sample, -1, 1)
1812+
if self._get_scheduler_config_value(scheduler, "clip_sample"):
1813+
clip_sample_range = self._get_scheduler_config_value(scheduler, "clip_sample_range", 1.0)
1814+
pred_original_sample = torch.clamp(pred_original_sample, -clip_sample_range, clip_sample_range)
17061815

17071816
# 4. Compute coefficients for pred_original_sample x_0 and current sample x_t
17081817
# See formula (7) from https://arxiv.org/pdf/2006.11239.pdf
@@ -1714,11 +1823,15 @@ def get_likelihood( # type: ignore[override]
17141823
predicted_mean = pred_original_sample_coeff * pred_original_sample + current_sample_coeff * noisy_image
17151824

17161825
# get the posterior mean and variance
1717-
posterior_mean = scheduler._get_mean(timestep=t, x_0=inputs, x_t=noisy_image) # type: ignore[operator]
1718-
posterior_variance = scheduler._get_variance(timestep=t, predicted_variance=predicted_variance) # type: ignore[operator]
1826+
posterior_mean = self._get_posterior_mean(scheduler=scheduler, timestep=t, x_0=inputs, x_t=noisy_image)
1827+
posterior_variance = self._get_posterior_variance(
1828+
scheduler=scheduler, timestep=t, predicted_variance=predicted_variance
1829+
)
17191830

17201831
log_posterior_variance = torch.log(posterior_variance)
1721-
log_predicted_variance = torch.log(predicted_variance) if predicted_variance else log_posterior_variance
1832+
log_predicted_variance = (
1833+
torch.log(predicted_variance) if predicted_variance is not None else log_posterior_variance
1834+
)
17221835

17231836
if t == 0:
17241837
# compute -log p(x_0|x_1)

‎tests/inferers/test_diffusion_inferer.py‎

Lines changed: 58 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424

2525
_, has_scipy = optional_import("scipy")
2626
_, has_einops = optional_import("einops")
27+
DiffusersDDPMScheduler, has_diffusers = optional_import("diffusers", name="DDPMScheduler")
2728

2829
TEST_CASES = [
2930
[
@@ -126,6 +127,63 @@ def test_ddpm_sampler(self, model_params, input_shape):
126127
)
127128
self.assertEqual(len(intermediates), 10)
128129

130+
@skipUnless(has_einops and has_diffusers, "Requires einops and diffusers")
131+
def test_diffusers_ddpm_call(self):
132+
device = "cuda:0" if torch.cuda.is_available() else "cpu"
133+
model = DiffusionModelUNet(
134+
spatial_dims=2,
135+
in_channels=1,
136+
out_channels=1,
137+
channels=[32, 64],
138+
attention_levels=[False, True],
139+
num_res_blocks=1,
140+
num_head_channels=32,
141+
)
142+
model.to(device)
143+
model.eval()
144+
scheduler = DiffusersDDPMScheduler(num_train_timesteps=1000, beta_schedule="linear", prediction_type="epsilon")
145+
scheduler.set_timesteps(num_inference_steps=50)
146+
inferer = DiffusionInferer(scheduler=scheduler)
147+
148+
batch_size = 2
149+
image_size = 32
150+
inputs = torch.randn(batch_size, 1, image_size, image_size).to(device)
151+
noise = torch.randn_like(inputs)
152+
timesteps = torch.randint(0, scheduler.config.num_train_timesteps, (batch_size,)).long().to(device)
153+
with torch.no_grad():
154+
prediction = inferer(inputs=inputs, diffusion_model=model, noise=noise, timesteps=timesteps)
155+
156+
self.assertEqual(prediction.shape, inputs.shape)
157+
scheduler.set_timesteps(num_inference_steps=2)
158+
sample = inferer.sample(input_noise=noise, diffusion_model=model, scheduler=scheduler, verbose=False)
159+
self.assertEqual(sample.shape, inputs.shape)
160+
161+
@skipUnless(has_einops and has_diffusers, "Requires einops and diffusers")
162+
def test_diffusers_ddpm_get_likelihood(self):
163+
device = "cuda:0" if torch.cuda.is_available() else "cpu"
164+
model = DiffusionModelUNet(
165+
spatial_dims=2,
166+
in_channels=1,
167+
out_channels=1,
168+
channels=[8],
169+
norm_num_groups=8,
170+
attention_levels=[True],
171+
num_res_blocks=1,
172+
num_head_channels=8,
173+
)
174+
model.to(device)
175+
model.eval()
176+
inputs = torch.randn(2, 1, 8, 8).to(device)
177+
scheduler = DiffusersDDPMScheduler(num_train_timesteps=10, beta_schedule="linear", prediction_type="epsilon")
178+
inferer = DiffusionInferer(scheduler=scheduler)
179+
scheduler.set_timesteps(num_inference_steps=10)
180+
likelihood, intermediates = inferer.get_likelihood(
181+
inputs=inputs, diffusion_model=model, scheduler=scheduler, save_intermediates=True
182+
)
183+
self.assertEqual(len(intermediates), 10)
184+
self.assertEqual(intermediates[0].shape, inputs.shape)
185+
self.assertEqual(likelihood.shape[0], inputs.shape[0])
186+
129187
@parameterized.expand(TEST_CASES)
130188
@skipUnless(has_einops, "Requires einops")
131189
def test_ddim_sampler(self, model_params, input_shape):

0 commit comments

Comments
 (0)