WanAnimate block swap fix
This commit is contained in:
@@ -4294,9 +4294,7 @@ class WanVideoSampler:
|
||||
bg_images = image_embeds.get("bg_images", None)
|
||||
|
||||
pose_input_latents = current_ref_images = face_images = None
|
||||
#if wananim_pose_latents is not None:
|
||||
#pose_input_latents = tensor_pingpong_pad(wananim_pose_latents, target_latent_len)
|
||||
#log.info(f"WanAnimate: Pose input {wananim_pose_latents.shape} padded to shape {pose_input_latents.shape}")
|
||||
|
||||
if wananim_face_pixels is not None:
|
||||
face_images = tensor_pingpong_pad(wananim_face_pixels, target_len)
|
||||
log.info(f"WanAnimate: Face input {wananim_face_pixels.shape} padded to shape {face_images.shape}")
|
||||
@@ -4307,11 +4305,6 @@ class WanVideoSampler:
|
||||
bg_images_in = tensor_pingpong_pad(bg_images, target_len)
|
||||
log.info(f"WanAnimate: BG images {bg_images.shape} padded to shape {bg_images.shape}")
|
||||
|
||||
# if replace_flag:
|
||||
# bg_images, mask_images = self.prepare_source_for_replace(src_bg_path, src_mask_path)
|
||||
# bg_images = inputs_padding(bg_images, target_len)
|
||||
# mask_images = inputs_padding(mask_images, target_len)
|
||||
|
||||
# init variables
|
||||
offloaded = False
|
||||
|
||||
@@ -4344,11 +4337,6 @@ class WanVideoSampler:
|
||||
vae.to(device)
|
||||
if ref_masks is not None:
|
||||
msk = ref_masks_in[:, start_latent:end_latent].to(device, dtype)
|
||||
# if msk.shape[1] < latent_window_size:
|
||||
# log.info(f"WanAnimate: Padding ref masks from {msk.shape} to length {latent_window_size}")
|
||||
# pad_length = latent_window_size - msk.shape[1]
|
||||
# last_frame = msk[:, -1:].repeat(1, pad_length, 1, 1)
|
||||
# msk = torch.cat([msk, last_frame], dim=1)
|
||||
else:
|
||||
msk = torch.zeros(4, latent_window_size, lat_h, lat_w, device=device, dtype=dtype)
|
||||
if bg_images is not None:
|
||||
|
||||
@@ -792,7 +792,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
vace_block_idx = None
|
||||
elif "blocks." in name:
|
||||
elif "blocks." in name and "face" not in name:
|
||||
try:
|
||||
block_idx = int(name.split("blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
@@ -831,7 +831,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
vace_block_idx = None
|
||||
elif "blocks." in name:
|
||||
elif "blocks." in name and "face" not in name:
|
||||
try:
|
||||
block_idx = int(name.split("blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
@@ -868,6 +868,8 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
cnt += 1
|
||||
if cnt % 100 == 0:
|
||||
pbar.update(100)
|
||||
#for name, param in transformer.named_parameters():
|
||||
# print(name, param.device, param.dtype)
|
||||
|
||||
pbar.update_absolute(0)
|
||||
|
||||
|
||||
@@ -2242,7 +2242,6 @@ class WanModel(torch.nn.Module):
|
||||
ip_img_ids[:, :, :, 2] = ip_img_ids[:, :, :, 2] + torch.linspace(w_len + freq_offset, w_len + freq_offset + w_ip - 1, steps=w_ip, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
|
||||
ip_img_ids = repeat(ip_img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
freqs_ip = self.rope_embedder(ip_img_ids).movedim(1, 2)
|
||||
#print("freqs_ip shape:", freqs_ip.shape)
|
||||
|
||||
# EchoShot cross attn freqs
|
||||
inner_c = None
|
||||
@@ -2256,9 +2255,6 @@ class WanModel(torch.nn.Module):
|
||||
x = torch.cat([x, motion_encoded], dim=1)
|
||||
freqs = torch.cat([freqs, freqs_motion], dim=1)
|
||||
|
||||
#t = torch.repeat_interleave(t, 2, dim=1)
|
||||
#t = torch.cat([t, torch.zeros((t.shape[0], 3), device=t.device, dtype=t.dtype)], dim=1)
|
||||
|
||||
# time embeddings
|
||||
if t.dim() == 2:
|
||||
b, f = t.shape
|
||||
|
||||
Reference in New Issue
Block a user