Allow using more than one start_image with the Fun InP models
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user