diff --git a/examples/scaffolded_scheduler/scheduling_ddpm_lite.py b/examples/scaffolded_scheduler/scheduling_ddpm_lite.py index 89f1003..df0b9e4 100644 --- a/examples/scaffolded_scheduler/scheduling_ddpm_lite.py +++ b/examples/scaffolded_scheduler/scheduling_ddpm_lite.py @@ -46,7 +46,8 @@ def set_timesteps(self, num_inference_steps: int, device: Union[str, torch.devic """ self.num_inference_steps = num_inference_steps step = self.config.num_train_timesteps // num_inference_steps - timesteps = (torch.arange(0, num_inference_steps) * step).round()[::-1].clone() + timesteps = (torch.arange(0, num_inference_steps) * step).round().long() + timesteps = torch.flip(timesteps, dims=[0]) # torch has no [::-1] self.timesteps = timesteps.to(device) if device is not None else timesteps def step( diff --git a/src/diffusers/schedulers/scheduling_euler_lite.py b/src/diffusers/schedulers/scheduling_euler_lite.py index c41edc5..5fc5795 100644 --- a/src/diffusers/schedulers/scheduling_euler_lite.py +++ b/src/diffusers/schedulers/scheduling_euler_lite.py @@ -40,7 +40,8 @@ def set_timesteps(self, num_inference_steps: int, device: Union[str, torch.devic """ self.num_inference_steps = num_inference_steps step = self.config.num_train_timesteps // num_inference_steps - timesteps = (torch.arange(0, num_inference_steps) * step).round()[::-1].clone() + timesteps = (torch.arange(0, num_inference_steps) * step).round().long() + timesteps = torch.flip(timesteps, dims=[0]) # torch has no [::-1] self.timesteps = timesteps.to(device) if device is not None else timesteps def step( diff --git a/templates/scheduler/scheduling_TEMPLATE.py b/templates/scheduler/scheduling_TEMPLATE.py index 907ed84..e52665a 100644 --- a/templates/scheduler/scheduling_TEMPLATE.py +++ b/templates/scheduler/scheduling_TEMPLATE.py @@ -36,7 +36,8 @@ def set_timesteps(self, num_inference_steps: int, device: Union[str, torch.devic """ self.num_inference_steps = num_inference_steps step = self.config.num_train_timesteps // num_inference_steps - timesteps = (torch.arange(0, num_inference_steps) * step).round()[::-1].clone() + timesteps = (torch.arange(0, num_inference_steps) * step).round().long() + timesteps = torch.flip(timesteps, dims=[0]) # torch has no [::-1] self.timesteps = timesteps.to(device) if device is not None else timesteps def step( diff --git a/tests/schedulers/test_scheduling_euler_lite.py b/tests/schedulers/test_scheduling_euler_lite.py index 8e4c2eb..3764f3b 100644 --- a/tests/schedulers/test_scheduling_euler_lite.py +++ b/tests/schedulers/test_scheduling_euler_lite.py @@ -103,6 +103,7 @@ def test_output_type(self): out = s.step(torch.ones_like(sample), 1, sample, generator=torch.Generator().manual_seed(0)) self.assertTrue(hasattr(out, "prev_sample")) self.assertEqual(out.prev_sample.shape, sample.shape) + self.assertEqual(out.prev_sample.dtype, sample.dtype) def test_same_seed_same_output(self): import torch