Allow using more than one start_image with the Fun InP models

This commit is contained in:
kijai
2025-03-30 23:06:28 +03:00
parent 4970f84f5f
commit facf65aec6
2 changed files with 8 additions and 8 deletions
+2 -2
View File
@@ -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)
+6 -6
View File
@@ -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)