diff --git a/src/common/diffusion/timesteps/sampling/trailing.py b/src/common/diffusion/timesteps/sampling/trailing.py index 955c8e5..a6524b0 100644 --- a/src/common/diffusion/timesteps/sampling/trailing.py +++ b/src/common/diffusion/timesteps/sampling/trailing.py @@ -36,7 +36,7 @@ class UniformTrailingSamplingTimesteps(SamplingTimesteps): dtype: torch.dtype = torch.float32, ): # Create trailing timesteps with specified dtype - timesteps = torch.arange(1.0, 0.0, -1.0 / steps, device=device, dtype=dtype) + timesteps = torch.arange(1.0, 0.0, -1.0 / steps, device='cpu').to(device=device, dtype=dtype) # Shift timesteps. timesteps = shift * timesteps / (1 + (shift - 1) * timesteps) diff --git a/src/data/image/transforms/area_resize.py b/src/data/image/transforms/area_resize.py index ead8188..5873b85 100644 --- a/src/data/image/transforms/area_resize.py +++ b/src/data/image/transforms/area_resize.py @@ -50,10 +50,12 @@ class AreaResize: resized_height, resized_width = round(height * scale), round(width * scale) + antialias = not (isinstance(image, torch.Tensor) and image.device.type == 'mps') return TVF.resize( image, size=(resized_height, resized_width), interpolation=self.interpolation, + antialias=antialias, ) diff --git a/src/data/image/transforms/side_resize.py b/src/data/image/transforms/side_resize.py index 6fff1f4..6d5273f 100644 --- a/src/data/image/transforms/side_resize.py +++ b/src/data/image/transforms/side_resize.py @@ -56,8 +56,9 @@ class SideResize: else: size = self.size - # Resize to shortest edge - resized = TVF.resize(image, size, self.interpolation) + # Resize to shortest edge (disable antialias only for MPS tensors - not supported) + antialias = not (isinstance(image, torch.Tensor) and image.device.type == 'mps') + resized = TVF.resize(image, size, self.interpolation, antialias=antialias) # Apply max_size constraint if specified if self.max_size > 0: @@ -69,6 +70,6 @@ class SideResize: if max(h, w) > self.max_size: scale = self.max_size / max(h, w) new_h, new_w = round(h * scale), round(w * scale) - resized = TVF.resize(resized, (new_h, new_w), self.interpolation) + resized = TVF.resize(resized, (new_h, new_w), self.interpolation, antialias=antialias) return resized