diff --git a/nodes.py b/nodes.py index c916ef2..0593d2d 100644 --- a/nodes.py +++ b/nodes.py @@ -1396,7 +1396,7 @@ class WanVideoImageToVideoEncode: mask = torch.zeros(1, base_frames, lat_h, lat_w, device=device) if start_image is not None: - mask[:, 0] = 1 # First frame + mask[:, 0:start_image.shape[0]] = 1 # First frame if end_image is not None and not fun_model: mask[:, -1] = 1 # End frame if exists @@ -1429,7 +1429,7 @@ class WanVideoImageToVideoEncode: vae.to(device) if start_image is not None and end_image is None: - zero_frames = torch.zeros(3, num_frames-1, H, W, device=device) + zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device) concatenated = torch.cat([resized_start_image.to(device), zero_frames], dim=1) elif start_image is None and end_image is not None: zero_frames = torch.zeros(3, num_frames-1, H, W, device=device) diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index ca966c8..12a6cd4 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -184,9 +184,9 @@ def attention( # ) attn_mask = None - q = q.transpose(1, 2).to(dtype) - k = k.transpose(1, 2).to(dtype) - v = v.transpose(1, 2).to(dtype) + q = q.transpose(1, 2)#.to(dtype) + k = k.transpose(1, 2)#.to(dtype) + v = v.transpose(1, 2)#.to(dtype) out = torch.nn.functional.scaled_dot_product_attention( q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p) @@ -196,9 +196,9 @@ def attention( elif attention_mode == 'sageattn': attn_mask = None - q = q.transpose(1, 2).to(dtype) - k = k.transpose(1, 2).to(dtype) - v = v.transpose(1, 2).to(dtype) + q = q.transpose(1, 2)#.to(dtype) + k = k.transpose(1, 2)#.to(dtype) + v = v.transpose(1, 2)#.to(dtype) out = sageattn_func( q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)