Fix MPS compatibility: disable antialias for MPS tensors, fix bfloat16 arange (#354)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user