diff --git a/src/diffusers/modular_pipelines/wan_animate_2/encoders.py b/src/diffusers/modular_pipelines/wan_animate_2/encoders.py index 21b70f636f7d..3eb85b41fe4c 100644 --- a/src/diffusers/modular_pipelines/wan_animate_2/encoders.py +++ b/src/diffusers/modular_pipelines/wan_animate_2/encoders.py @@ -81,11 +81,19 @@ def clip_visual_encode(image_encoder, tensor, device, dtype): return out.hidden_states[-2] -def get_i2v_mask(lat_t, lat_h, lat_w, mask_len=1, device="cuda"): +def get_i2v_mask(lat_t, lat_h, lat_w, mask_len=1, device=None): """Create an i2v mask in latent space. mask_len is in PIXEL space. Returns [4, lat_t, lat_h, lat_w] (no batch dim). + + Args: + device: device on which the mask is allocated. Must be passed explicitly by the caller. """ + if device is None: + raise ValueError( + "`device` must be specified when calling `get_i2v_mask`. It used to default to 'cuda', which " + "silently allocated the mask on CUDA and broke every non-CUDA accelerator (NPU/XPU/MPS/CPU)." + ) msk = torch.zeros(1, (lat_t - 1) * 4 + 1, lat_h, lat_w, device=device) msk[:, :mask_len] = 1 msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1) diff --git a/src/diffusers/pipelines/wan/pipeline_wan_animate.py b/src/diffusers/pipelines/wan/pipeline_wan_animate.py index a923219a7550..c2416243f111 100644 --- a/src/diffusers/pipelines/wan/pipeline_wan_animate.py +++ b/src/diffusers/pipelines/wan/pipeline_wan_animate.py @@ -465,8 +465,10 @@ def get_i2v_mask( mask_len: int = 1, mask_pixel_values: torch.Tensor | None = None, dtype: torch.dtype | None = None, - device: str | torch.device = "cuda", + device: str | torch.device | None = None, ) -> torch.Tensor: + device = device or self._execution_device + # mask_pixel_values shape (if supplied): [B, C = 1, T, latent_h, latent_w] if mask_pixel_values is None: mask_lat_size = torch.zeros(