Fix MPS compatibility: disable antialias for MPS tensors, fix bfloat16 arange (#354)

This commit is contained in:
Adrien Toupet
2025-12-03 13:09:59 -05:00
parent ed53581359
commit b40f26167c
3 changed files with 7 additions and 4 deletions
@@ -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)
+2
View File
@@ -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,
)
+4 -3
View File
@@ -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