diff --git a/src/diffusers/schedulers/scheduling_unipc_multistep.py b/src/diffusers/schedulers/scheduling_unipc_multistep.py index 5c2cbcc13ff1..e8f5af16fbe4 100644 --- a/src/diffusers/schedulers/scheduling_unipc_multistep.py +++ b/src/diffusers/schedulers/scheduling_unipc_multistep.py @@ -903,7 +903,7 @@ def multistep_uni_p_bh_update( rks.append(rk) D1s.append((mi - m0) / rk) - rks.append(torch.ones((), device=device)) + rks.append(torch.ones_like(h)) rks = torch.stack(rks) R = [] @@ -1038,7 +1038,7 @@ def multistep_uni_c_bh_update( rks.append(rk) D1s.append((mi - m0) / rk) - rks.append(torch.ones((), device=device)) + rks.append(torch.ones_like(h)) rks = torch.stack(rks) R = [] diff --git a/tests/schedulers/test_scheduler_unipc.py b/tests/schedulers/test_scheduler_unipc.py index ac7e1d3f88b4..9f5c37667e58 100644 --- a/tests/schedulers/test_scheduler_unipc.py +++ b/tests/schedulers/test_scheduler_unipc.py @@ -257,6 +257,20 @@ def test_fp16_support(self): assert sample.dtype == torch.float16 + def test_default_dtype_float64(self): + default_dtype = torch.get_default_dtype() + torch.set_default_dtype(torch.float64) + try: + sample = self.full_loop(solver_order=3) + finally: + torch.set_default_dtype(default_dtype) + + result_mean = torch.mean(torch.abs(sample)) + reference_mean = torch.mean(torch.abs(self.full_loop(solver_order=3))) + + assert sample.dtype == torch.float64 + assert abs(result_mean.item() - reference_mean.item()) < 1e-3 + def test_full_loop_with_noise(self): scheduler_class = self.scheduler_classes[0] scheduler_config = self.get_scheduler_config()