WanAnimate: Fix first latent issue when not using bg_images and looping
This commit is contained in:
+1
-1
@@ -219,7 +219,7 @@ class HuMoEmbeds:
|
||||
vae.to(device)
|
||||
zero_frames = torch.zeros(1, 3, pixel_frame_num + 4*num_refs, height, width, device=device, dtype=vae.dtype)
|
||||
zero_latents = vae.encode(zero_frames, device=device, tiled=tiled_vae)[0].to(offload_device)
|
||||
vae.model.clear_cache()
|
||||
|
||||
vae.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
|
||||
|
||||
@@ -1018,7 +1018,6 @@ class WanVideoImageToVideoEncode:
|
||||
|
||||
vae.to(device)
|
||||
y = vae.encode([concatenated], device, end_=(end_image is not None and not fun_or_fl2v_model),tiled=tiled_vae)[0]
|
||||
vae.model.clear_cache()
|
||||
del concatenated
|
||||
|
||||
has_ref = False
|
||||
@@ -1141,7 +1140,7 @@ class WanVideoAnimateEmbeds:
|
||||
resized_pose_images = resized_pose_images * 2 - 1
|
||||
pose_latents = vae.encode([resized_pose_images.to(device, vae.dtype)], device,tiled=tiled_vae)
|
||||
pose_latents = pose_latents.to(offload_device)
|
||||
vae.model.clear_cache()
|
||||
|
||||
if not looping and pose_latents.shape[2] < latent_window_size:
|
||||
log.info(f"WanAnimate: Padding pose latents from {pose_latents.shape} to length {latent_window_size}")
|
||||
pad_len = latent_window_size - pose_latents.shape[2]
|
||||
@@ -1156,15 +1155,15 @@ class WanVideoAnimateEmbeds:
|
||||
resized_bg_images = common_upscale(bg_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
|
||||
else:
|
||||
resized_bg_images = bg_images.permute(3, 0, 1, 2) # C, T, H, W
|
||||
resized_bg_images = resized_bg_images[:3] * 2 - 1
|
||||
if not looping:
|
||||
bg_latents = vae.encode([resized_bg_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0]
|
||||
bg_latents = bg_latents.to(offload_device)
|
||||
vae.model.clear_cache()
|
||||
print("bg_latents", bg_latents.shape)
|
||||
del resized_bg_images
|
||||
else:
|
||||
resized_bg_images = resized_bg_images.to(offload_device, dtype=vae.dtype)
|
||||
resized_bg_images = (resized_bg_images[:3] * 2 - 1)
|
||||
|
||||
if not looping:
|
||||
if bg_images is None:
|
||||
resized_bg_images = torch.zeros(3, num_frames - num_refs, H, W, device=device, dtype=vae.dtype)
|
||||
bg_latents = vae.encode([resized_bg_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0].to(offload_device)
|
||||
del resized_bg_images
|
||||
elif bg_images is not None:
|
||||
resized_bg_images = resized_bg_images.to(offload_device, dtype=vae.dtype)
|
||||
|
||||
if ref_images is not None:
|
||||
if ref_images.shape[1] != H or ref_images.shape[2] != W:
|
||||
@@ -1173,46 +1172,31 @@ class WanVideoAnimateEmbeds:
|
||||
resized_ref_images = ref_images.permute(3, 0, 1, 2) # C, T, H, W
|
||||
resized_ref_images = resized_ref_images[:3] * 2 - 1
|
||||
|
||||
if looping or bg_images is not None: # looping or when using background, encode refs separately
|
||||
ref_latent = vae.encode([resized_ref_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0]
|
||||
msk = torch.zeros(4, 1, lat_h, lat_w, device=device, dtype=vae.dtype)
|
||||
msk[:, :1] = 1
|
||||
ref_latent_masked = torch.cat([msk, ref_latent], dim=0) # 4+C 1 H W
|
||||
ref_latent_masked = ref_latent_masked.to(offload_device)
|
||||
|
||||
if bg_images is None:
|
||||
zero_frames = torch.zeros(3, num_frames - num_refs, H, W, device=device, dtype=vae.dtype)
|
||||
concatenated = torch.cat([resized_ref_images.to(device, dtype=vae.dtype), zero_frames], dim=1)
|
||||
del zero_frames
|
||||
ref_latent = vae.encode([concatenated.to(device, vae.dtype)], device,tiled=tiled_vae)[0]
|
||||
ref_latent = ref_latent.to(offload_device)
|
||||
del concatenated
|
||||
|
||||
vae.model.clear_cache()
|
||||
ref_latent = vae.encode([resized_ref_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0]
|
||||
msk = torch.zeros(4, 1, lat_h, lat_w, device=device, dtype=vae.dtype)
|
||||
msk[:, :num_refs] = 1
|
||||
ref_latent_masked = torch.cat([msk, ref_latent], dim=0).to(offload_device) # 4+C 1 H W
|
||||
|
||||
if mask is None:
|
||||
ref_mask = torch.zeros(1, num_frames, lat_h, lat_w, device=offload_device, dtype=vae.dtype)
|
||||
bg_mask = torch.zeros(1, num_frames, lat_h, lat_w, device=offload_device, dtype=vae.dtype)
|
||||
else:
|
||||
ref_mask = 1 - mask[:num_frames]
|
||||
if ref_mask.shape[0] < num_frames and not looping:
|
||||
ref_mask = torch.cat([ref_mask, ref_mask[-1:].repeat(num_frames - ref_mask.shape[0], 1, 1)], dim=0)
|
||||
ref_mask = common_upscale(ref_mask.unsqueeze(1), lat_w, lat_h, "nearest", "disabled").squeeze(1)
|
||||
ref_mask = ref_mask.to(vae.dtype).to(offload_device)
|
||||
ref_mask = ref_mask.unsqueeze(-1).permute(3, 0, 1, 2) # C, T, H, W
|
||||
bg_mask = 1 - mask[:num_frames]
|
||||
if bg_mask.shape[0] < num_frames and not looping:
|
||||
bg_mask = torch.cat([bg_mask, bg_mask[-1:].repeat(num_frames - bg_mask.shape[0], 1, 1)], dim=0)
|
||||
bg_mask = common_upscale(bg_mask.unsqueeze(1), lat_w, lat_h, "nearest", "disabled").squeeze(1)
|
||||
bg_mask = bg_mask.unsqueeze(-1).permute(3, 0, 1, 2).to(offload_device, vae.dtype) # C, T, H, W
|
||||
|
||||
if bg_images is None:
|
||||
ref_mask[:, :num_refs] = 1
|
||||
ref_mask_mask_repeated = torch.repeat_interleave(ref_mask[:, 0:1], repeats=4, dim=1) # T, C, H, W
|
||||
ref_mask = torch.cat([ref_mask_mask_repeated, ref_mask[:, 1:]], dim=1)
|
||||
ref_mask = ref_mask.view(1, ref_mask.shape[1] // 4, 4, lat_h, lat_w) # 1, T, C, H, W
|
||||
ref_mask = ref_mask.movedim(1, 2)[0]# C, T, H, W
|
||||
if bg_images is None and looping:
|
||||
bg_mask[:, :num_refs] = 1
|
||||
bg_mask_mask_repeated = torch.repeat_interleave(bg_mask[:, 0:1], repeats=4, dim=1) # T, C, H, W
|
||||
bg_mask = torch.cat([bg_mask_mask_repeated, bg_mask[:, 1:]], dim=1)
|
||||
bg_mask = bg_mask.view(1, bg_mask.shape[1] // 4, 4, lat_h, lat_w) # 1, T, C, H, W
|
||||
bg_mask = bg_mask.movedim(1, 2)[0]# C, T, H, W
|
||||
|
||||
if not looping:
|
||||
if bg_images is not None:
|
||||
bg_latents_masked = torch.cat([ref_mask[:, :bg_latents.shape[1]], bg_latents], dim=0)
|
||||
ref_latent = torch.cat([ref_latent_masked, bg_latents_masked], dim=1)
|
||||
else:
|
||||
ref_latent = torch.cat([ref_mask, ref_latent], dim=0)
|
||||
bg_latents_masked = torch.cat([bg_mask[:, :bg_latents.shape[1]], bg_latents], dim=0)
|
||||
del bg_mask, bg_latents
|
||||
ref_latent = torch.cat([ref_latent_masked, bg_latents_masked], dim=1)
|
||||
else:
|
||||
ref_latent = ref_latent_masked
|
||||
|
||||
@@ -1225,7 +1209,6 @@ class WanVideoAnimateEmbeds:
|
||||
resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0)
|
||||
resized_face_images = resized_face_images.to(offload_device, dtype=vae.dtype)
|
||||
|
||||
vae.model.clear_cache()
|
||||
|
||||
seq_len = math.ceil((target_shape[2] * target_shape[3]) / 4 * target_shape[1])
|
||||
|
||||
@@ -1240,7 +1223,7 @@ class WanVideoAnimateEmbeds:
|
||||
"max_seq_len": seq_len,
|
||||
"pose_latents": pose_latents,
|
||||
"bg_images": resized_bg_images if bg_images is not None and looping else None,
|
||||
"ref_masks": ref_mask if mask is not None and looping else None,
|
||||
"ref_masks": bg_mask if mask is not None and looping else None,
|
||||
"ref_latent": ref_latent,
|
||||
"ref_image": resized_ref_images if ref_images is not None else None,
|
||||
"face_pixels": resized_face_images if face_images is not None else None,
|
||||
@@ -1625,7 +1608,7 @@ class WanVideoVACEEncode:
|
||||
|
||||
vae = vae.to(device)
|
||||
z0 = self.vace_encode_frames(vae, input_frames, ref_images, masks=input_masks, tiled_vae=tiled_vae)
|
||||
vae.model.clear_cache()
|
||||
|
||||
m0 = self.vace_encode_masks(input_masks, ref_images)
|
||||
z = self.vace_latent(z0, m0)
|
||||
vae.to(offload_device)
|
||||
@@ -1666,7 +1649,7 @@ class WanVideoVACEEncode:
|
||||
reactive = vae.encode(reactive, device=device, tiled=tiled_vae)
|
||||
latents = [torch.cat((u, c), dim=0) for u, c in zip(inactive, reactive)]
|
||||
del inactive, reactive
|
||||
vae.model.clear_cache()
|
||||
|
||||
|
||||
cat_latents = []
|
||||
for latent, refs in zip(latents, ref_images):
|
||||
@@ -3858,7 +3841,7 @@ class WanVideoSampler:
|
||||
else:
|
||||
cond_ = cond_image if is_first_clip else cond_frame
|
||||
latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
|
||||
vae.model.clear_cache()
|
||||
|
||||
vae.to(offload_device)
|
||||
|
||||
#motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1
|
||||
@@ -3957,7 +3940,7 @@ class WanVideoSampler:
|
||||
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
|
||||
device=device, tiled=tiled_vae
|
||||
).to(dtype)
|
||||
vae.model.clear_cache()
|
||||
|
||||
vae.to(offload_device)
|
||||
uni3c_data['render_latent'] = render_latent
|
||||
|
||||
@@ -4046,7 +4029,7 @@ class WanVideoSampler:
|
||||
latent = latent[:,:-humo_reference_count]
|
||||
vae.to(device)
|
||||
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
|
||||
vae.model.clear_cache()
|
||||
|
||||
vae.to(offload_device)
|
||||
|
||||
sampling_pbar.close()
|
||||
@@ -4190,13 +4173,13 @@ class WanVideoSampler:
|
||||
log.info(f"Sampling {total_frames} frames in {s2v_num_repeat} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
|
||||
# sample
|
||||
for r in range(s2v_num_repeat):
|
||||
vae.model.clear_cache()
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
if ref_motion_image is not None:
|
||||
vae.to(device)
|
||||
ref_motion = vae.encode(ref_motion_image.to(vae.dtype), device=device, pbar=False).to(dtype)[0]
|
||||
vae.model.clear_cache()
|
||||
|
||||
vae.to(offload_device)
|
||||
|
||||
left_idx = r * infer_frames
|
||||
@@ -4265,7 +4248,7 @@ class WanVideoSampler:
|
||||
ref_motion_image = videos_last_frames
|
||||
|
||||
vae.to(offload_device)
|
||||
vae.model.clear_cache()
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gen_video_samples = torch.cat(framepack_out, dim=2).squeeze(0).permute(1, 2, 3, 0)
|
||||
|
||||
@@ -4300,7 +4283,7 @@ class WanVideoSampler:
|
||||
ref_masks = image_embeds.get("ref_masks", None)
|
||||
bg_images = image_embeds.get("bg_images", None)
|
||||
|
||||
pose_input_latents = current_ref_images = face_images = None
|
||||
current_ref_images = face_images = None
|
||||
|
||||
if wananim_face_pixels is not None:
|
||||
face_images = tensor_pingpong_pad(wananim_face_pixels, target_len)
|
||||
@@ -4365,7 +4348,7 @@ class WanVideoSampler:
|
||||
else:
|
||||
temporal_ref_latents = temporal_ref_latents[:, :msk.shape[1]]
|
||||
|
||||
vae.model.clear_cache()
|
||||
|
||||
vae.to(offload_device)
|
||||
|
||||
temporal_ref_latents = torch.cat([msk, temporal_ref_latents], dim=0) # 4+C T H W
|
||||
@@ -4501,7 +4484,7 @@ class WanVideoSampler:
|
||||
vae.to(device)
|
||||
videos = vae.decode(latent[:, 1:].unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
|
||||
del latent
|
||||
vae.model.clear_cache()
|
||||
|
||||
|
||||
sampling_pbar.close()
|
||||
|
||||
@@ -4778,7 +4761,7 @@ class WanVideoDecode:
|
||||
if end_image is not None:
|
||||
enable_vae_tiling = False
|
||||
images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0]
|
||||
vae.model.clear_cache()
|
||||
|
||||
|
||||
images = images.cpu().float()
|
||||
|
||||
@@ -4798,7 +4781,7 @@ class WanVideoDecode:
|
||||
if end_image is not None:
|
||||
images = images[:, 0:-1]
|
||||
|
||||
vae.model.clear_cache()
|
||||
|
||||
vae.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
|
||||
@@ -4849,7 +4832,7 @@ class WanVideoEncodeLatentBatch:
|
||||
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))
|
||||
else:
|
||||
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling)
|
||||
vae.model.clear_cache()
|
||||
|
||||
if latent_strength != 1.0:
|
||||
latent *= latent_strength
|
||||
latent_list.append(latent.squeeze(0).cpu())
|
||||
@@ -4909,7 +4892,7 @@ class WanVideoEncode:
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
latents = vae.encode(image * 2.0 - 1.0, device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))
|
||||
vae.model.clear_cache()
|
||||
|
||||
vae.to(offload_device)
|
||||
if latent_strength != 1.0:
|
||||
latents *= latent_strength
|
||||
|
||||
@@ -124,7 +124,6 @@ class WanVideoImageClipEncode:
|
||||
|
||||
y = torch.concat([mask, y])
|
||||
|
||||
vae.model.clear_cache()
|
||||
vae.to(offload_device)
|
||||
|
||||
image_embeds = {
|
||||
|
||||
@@ -17,20 +17,6 @@ def check_is_instance(model, module_class):
|
||||
return False
|
||||
|
||||
|
||||
def block_causal_mask(x, block_size):
|
||||
# params
|
||||
b, n, s, _, device = *x.size(), x.device
|
||||
assert s % block_size == 0
|
||||
num_blocks = s // block_size
|
||||
|
||||
# build mask
|
||||
mask = torch.zeros(b, n, s, s, dtype=torch.bool, device=device)
|
||||
for i in range(num_blocks):
|
||||
mask[:, :,
|
||||
i * block_size:(i + 1) * block_size, :(i + 1) * block_size] = 1
|
||||
return mask
|
||||
|
||||
|
||||
class CausalConv3d(nn.Conv3d):
|
||||
"""
|
||||
Causal 3d convolusion.
|
||||
@@ -329,12 +315,7 @@ class AttentionBlock(nn.Module):
|
||||
0, 1, 3, 2).contiguous().chunk(3, dim=-1)
|
||||
|
||||
# apply attention
|
||||
x = F.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
#attn_mask=block_causal_mask(q, block_size=h * w)
|
||||
)
|
||||
x = F.scaled_dot_product_attention(q, k, v)
|
||||
x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w)
|
||||
|
||||
# output
|
||||
@@ -860,7 +841,7 @@ class Decoder3d_38(nn.Module):
|
||||
|
||||
# dimensions
|
||||
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
|
||||
scale = 1.0 / 2 ** (len(dim_mult) - 2)
|
||||
|
||||
# init block
|
||||
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
|
||||
|
||||
@@ -989,8 +970,6 @@ class VideoVAE_(nn.Module):
|
||||
|
||||
#modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames
|
||||
def encode_2(self, x):
|
||||
self.clear_cache()
|
||||
## cache
|
||||
t = x.shape[2]
|
||||
iter_ = 2 + (t - 2) // 4
|
||||
|
||||
@@ -1010,18 +989,16 @@ class VideoVAE_(nn.Module):
|
||||
feat_cache=self._enc_feat_map,
|
||||
feat_idx=self._enc_conv_idx)
|
||||
out = torch.cat([out, out_], 2)
|
||||
self.clear_cache()
|
||||
out_head = out[:, :, :iter_ - 1, :, :]
|
||||
out_tail = out[:, :, -1, :, :].unsqueeze(2)
|
||||
mu = torch.cat([self.conv1(out_head), self.conv1(out_tail)], dim=2).chunk(2, dim=1)[0]
|
||||
|
||||
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
|
||||
|
||||
return mu
|
||||
|
||||
|
||||
def encode(self, x, pbar=True):
|
||||
self.clear_cache()
|
||||
## cache
|
||||
t = x.shape[2]
|
||||
iter_ = 1 + (t - 1) // 4
|
||||
if pbar:
|
||||
@@ -1040,18 +1017,17 @@ class VideoVAE_(nn.Module):
|
||||
out = torch.cat([out, out_], 2)
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
self.clear_cache()
|
||||
mu = self.conv1(out).chunk(2, dim=1)[0]
|
||||
|
||||
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
|
||||
if pbar:
|
||||
pbar.update_absolute(0)
|
||||
|
||||
|
||||
return mu
|
||||
|
||||
|
||||
#modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames
|
||||
def decode_2(self, z):
|
||||
self.clear_cache()
|
||||
# z: [b,c,t,h,w]
|
||||
|
||||
z = z / self.inv_std.to(z) + self.mean.to(z)
|
||||
@@ -1076,12 +1052,12 @@ class VideoVAE_(nn.Module):
|
||||
feat_cache=self._feat_map,
|
||||
feat_idx=self._conv_idx)
|
||||
out = torch.cat([out, out_], 2) # may add tensor offload
|
||||
self.clear_cache()
|
||||
return out
|
||||
|
||||
|
||||
|
||||
def decode(self, z, pbar=True):
|
||||
self.clear_cache()
|
||||
# z: [b,c,t,h,w]
|
||||
z = z / self.inv_std.to(z) + self.mean.to(z)
|
||||
iter_ = z.shape[2]
|
||||
@@ -1104,6 +1080,7 @@ class VideoVAE_(nn.Module):
|
||||
pbar.update(1)
|
||||
if pbar:
|
||||
pbar.update_absolute(0)
|
||||
self.clear_cache()
|
||||
return out
|
||||
|
||||
def reparameterize(self, mu, log_var):
|
||||
@@ -1318,6 +1295,7 @@ class WanVideoVAE(nn.Module):
|
||||
return video
|
||||
|
||||
def encode(self, videos, device, tiled=False,end_=False, tile_size=None, tile_stride=None, pbar=True):
|
||||
self.model.clear_cache()
|
||||
videos = [video.to("cpu") for video in videos]
|
||||
hidden_states = []
|
||||
for video in videos:
|
||||
@@ -1336,6 +1314,7 @@ class WanVideoVAE(nn.Module):
|
||||
|
||||
|
||||
def decode(self, hidden_states, device, tiled=False, end_=False, tile_size=(34, 34), tile_stride=(18, 16), pbar=True):
|
||||
self.model.clear_cache()
|
||||
hidden_states = [hidden_state.to("cpu") for hidden_state in hidden_states]
|
||||
videos = []
|
||||
for hidden_state in hidden_states:
|
||||
@@ -1352,25 +1331,6 @@ class WanVideoVAE(nn.Module):
|
||||
return videos
|
||||
|
||||
|
||||
@staticmethod
|
||||
def state_dict_converter():
|
||||
return WanVideoVAEStateDictConverter()
|
||||
|
||||
|
||||
class WanVideoVAEStateDictConverter:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def from_civitai(self, state_dict):
|
||||
state_dict_ = {}
|
||||
if 'model_state' in state_dict:
|
||||
state_dict = state_dict['model_state']
|
||||
for name in state_dict:
|
||||
state_dict_['model.' + name] = state_dict[name]
|
||||
return state_dict_
|
||||
|
||||
|
||||
class VideoVAE38_(VideoVAE_):
|
||||
|
||||
def __init__(self,
|
||||
@@ -1427,9 +1387,7 @@ class VideoVAE38_(VideoVAE_):
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
mu = self.conv1(out).chunk(2, dim=1)[0]
|
||||
|
||||
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
|
||||
|
||||
self.clear_cache()
|
||||
return mu
|
||||
|
||||
|
||||
Reference in New Issue
Block a user