From 555a91e14a89391813b75ddd64e1928f1d475bf4 Mon Sep 17 00:00:00 2001 From: lllakshit Date: Sun, 23 Aug 2026 23:44:13 +0530 Subject: [PATCH] Fix NaN samples from DDPM fixed_large_log variance _get_variance(fixed_large_log) returns log(beta). step() treated that value like a variance and took sqrt(), which is NaN for beta < 1. Convert log variance with exp(0.5 * log) for both DDPMScheduler and DDPMParallelScheduler, matching learned_range and the fixed_large scale. Fixes #14569 --- src/diffusers/schedulers/scheduling_ddpm.py | 5 ++++- .../schedulers/scheduling_ddpm_parallel.py | 5 ++++- tests/schedulers/test_scheduler_ddpm.py | 18 ++++++++++++++++++ .../schedulers/test_scheduler_ddpm_parallel.py | 18 ++++++++++++++++++ 4 files changed, 44 insertions(+), 2 deletions(-) diff --git a/src/diffusers/schedulers/scheduling_ddpm.py b/src/diffusers/schedulers/scheduling_ddpm.py index 972c46c6e930..5f83722a4c79 100644 --- a/src/diffusers/schedulers/scheduling_ddpm.py +++ b/src/diffusers/schedulers/scheduling_ddpm.py @@ -550,7 +550,10 @@ def step( ) if self.variance_type == "fixed_small_log": variance = self._get_variance(t, predicted_variance=predicted_variance) * variance_noise - elif self.variance_type == "learned_range": + elif self.variance_type in ("learned_range", "fixed_large_log"): + # `_get_variance` returns a log-space value for these types. Convert back to a + # standard deviation with exp(0.5 * log_var). Taking sqrt() of log(beta) for + # `fixed_large_log` is invalid (beta < 1) and produces NaNs. variance = self._get_variance(t, predicted_variance=predicted_variance) variance = torch.exp(0.5 * variance) * variance_noise else: diff --git a/src/diffusers/schedulers/scheduling_ddpm_parallel.py b/src/diffusers/schedulers/scheduling_ddpm_parallel.py index 725958b01a44..40c0fff8c67e 100644 --- a/src/diffusers/schedulers/scheduling_ddpm_parallel.py +++ b/src/diffusers/schedulers/scheduling_ddpm_parallel.py @@ -565,7 +565,10 @@ def step( ) if self.variance_type == "fixed_small_log": variance = self._get_variance(t, predicted_variance=predicted_variance) * variance_noise - elif self.variance_type == "learned_range": + elif self.variance_type in ("learned_range", "fixed_large_log"): + # `_get_variance` returns a log-space value for these types. Convert back to a + # standard deviation with exp(0.5 * log_var). Taking sqrt() of log(beta) for + # `fixed_large_log` is invalid (beta < 1) and produces NaNs. variance = self._get_variance(t, predicted_variance=predicted_variance) variance = torch.exp(0.5 * variance) * variance_noise else: diff --git a/tests/schedulers/test_scheduler_ddpm.py b/tests/schedulers/test_scheduler_ddpm.py index 056b5d83350e..5c76cd002bca 100644 --- a/tests/schedulers/test_scheduler_ddpm.py +++ b/tests/schedulers/test_scheduler_ddpm.py @@ -68,6 +68,24 @@ def test_variance(self): assert torch.sum(torch.abs(scheduler._get_variance(487) - 0.00979)) < 1e-5 assert torch.sum(torch.abs(scheduler._get_variance(999) - 0.02)) < 1e-5 + def test_fixed_large_log_sampling_is_finite_and_matches_fixed_large(self): + # `fixed_large_log` stores log(beta). Sampling must use exp(0.5 * log) rather than + # sqrt(log(beta)), which is NaN because beta < 1. The intended scale matches `fixed_large`. + sample = torch.zeros((1, 2, 2, 2)) + model_output = torch.zeros_like(sample) + log_scheduler = self.scheduler_classes[0](**self.get_scheduler_config(variance_type="fixed_large_log")) + large_scheduler = self.scheduler_classes[0](**self.get_scheduler_config(variance_type="fixed_large")) + + log_output = log_scheduler.step( + model_output, 500, sample, generator=torch.Generator().manual_seed(0) + ).prev_sample + large_output = large_scheduler.step( + model_output, 500, sample, generator=torch.Generator().manual_seed(0) + ).prev_sample + + assert torch.isfinite(log_output).all() + assert torch.allclose(log_output, large_output) + def test_rescale_betas_zero_snr(self): for rescale_betas_zero_snr in [True, False]: self.check_over_configs(rescale_betas_zero_snr=rescale_betas_zero_snr) diff --git a/tests/schedulers/test_scheduler_ddpm_parallel.py b/tests/schedulers/test_scheduler_ddpm_parallel.py index 377067071c25..1655bf42f08b 100644 --- a/tests/schedulers/test_scheduler_ddpm_parallel.py +++ b/tests/schedulers/test_scheduler_ddpm_parallel.py @@ -82,6 +82,24 @@ def test_variance(self): assert torch.sum(torch.abs(scheduler._get_variance(487) - 0.00979)) < 1e-5 assert torch.sum(torch.abs(scheduler._get_variance(999) - 0.02)) < 1e-5 + def test_fixed_large_log_sampling_is_finite_and_matches_fixed_large(self): + # `fixed_large_log` stores log(beta). Sampling must use exp(0.5 * log) rather than + # sqrt(log(beta)), which is NaN because beta < 1. The intended scale matches `fixed_large`. + sample = torch.zeros((1, 2, 2, 2)) + model_output = torch.zeros_like(sample) + log_scheduler = self.scheduler_classes[0](**self.get_scheduler_config(variance_type="fixed_large_log")) + large_scheduler = self.scheduler_classes[0](**self.get_scheduler_config(variance_type="fixed_large")) + + log_output = log_scheduler.step( + model_output, 500, sample, generator=torch.Generator().manual_seed(0) + ).prev_sample + large_output = large_scheduler.step( + model_output, 500, sample, generator=torch.Generator().manual_seed(0) + ).prev_sample + + assert torch.isfinite(log_output).all() + assert torch.allclose(log_output, large_output) + def test_rescale_betas_zero_snr(self): for rescale_betas_zero_snr in [True, False]: self.check_over_configs(rescale_betas_zero_snr=rescale_betas_zero_snr)