Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 9 additions & 1 deletion src/diffusers/modular_pipelines/wan_animate_2/encoders.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
4 changes: 3 additions & 1 deletion src/diffusers/pipelines/wan/pipeline_wan_animate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading