Support WanAnimate
This commit is contained in:
+1
-1
@@ -13,7 +13,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
|
||||
module_prefix = prefix + name + "."
|
||||
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights)
|
||||
|
||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix:
|
||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix and "face" not in module_prefix:
|
||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
||||
out_features = state_dict[module_prefix + "weight"].shape[0]
|
||||
if scale_weights is not None:
|
||||
|
||||
@@ -80,6 +80,39 @@ def offload_transformer(transformer):
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
def init_blockswap(transformer, block_swap_args, model):
|
||||
if not transformer.patched_linear:
|
||||
if block_swap_args is not None:
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name:
|
||||
param.data = param.data.to(device)
|
||||
if "control_adapter" in name:
|
||||
param.data = param.data.to(device)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
block_swap_args["offload_txt_emb"],
|
||||
block_swap_args["offload_img_emb"],
|
||||
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
|
||||
)
|
||||
elif model["auto_cpu_offload"]:
|
||||
for module in transformer.modules():
|
||||
if hasattr(module, "offload"):
|
||||
module.offload()
|
||||
if hasattr(module, "onload"):
|
||||
module.onload()
|
||||
for block in transformer.blocks:
|
||||
block.modulation = torch.nn.Parameter(block.modulation.to(device))
|
||||
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
|
||||
else:
|
||||
transformer.to(device)
|
||||
|
||||
|
||||
class WanVideoEnhanceAVideo:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -1031,6 +1064,190 @@ class WanVideoImageToVideoEncode:
|
||||
|
||||
return (image_embeds,)
|
||||
|
||||
# region WanAnimate
|
||||
class WanVideoAnimateEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"vae": ("WANVAE",),
|
||||
"width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}),
|
||||
"height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}),
|
||||
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
|
||||
"force_offload": ("BOOLEAN", {"default": True}),
|
||||
"frame_window_size": ("INT", {"default": 77, "min": 1, "max": 1000, "step": 1, "tooltip": "Number of frames to use for temporal attention window"}),
|
||||
"colormatch": (
|
||||
[
|
||||
'disabled',
|
||||
'mkl',
|
||||
'hm',
|
||||
'reinhard',
|
||||
'mvgd',
|
||||
'hm-mvgd-hm',
|
||||
'hm-mkl-hm',
|
||||
], {
|
||||
"default": 'disabled', "tooltip": "Color matching method to use between the windows"
|
||||
},),
|
||||
"pose_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional multiplier for the pose"}),
|
||||
"face_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional multiplier for the face"}),
|
||||
},
|
||||
"optional": {
|
||||
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
|
||||
"ref_images": ("IMAGE", {"tooltip": "Image to encode"}),
|
||||
"pose_images": ("IMAGE", {"tooltip": "end frame"}),
|
||||
"face_images": ("IMAGE", {"tooltip": "end frame"}),
|
||||
"bg_images": ("IMAGE", {"tooltip": "background images"}),
|
||||
"mask": ("MASK", {"tooltip": "mask"}),
|
||||
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, vae, width, height, num_frames, force_offload, frame_window_size, colormatch, pose_strength, face_strength,
|
||||
ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None):
|
||||
|
||||
H = height
|
||||
W = width
|
||||
|
||||
lat_h = H // vae.upsampling_factor
|
||||
lat_w = W // vae.upsampling_factor
|
||||
|
||||
num_refs = ref_images.shape[0] if ref_images is not None else 0
|
||||
|
||||
num_frames = ((num_frames - 1) // 4) * 4 + 1
|
||||
target_shape = (16, (num_frames - 1) // 4 + 1 + num_refs, lat_h, lat_w)
|
||||
latent_window_size = ((frame_window_size - 1) // 4) + 1
|
||||
|
||||
looping = num_frames > frame_window_size
|
||||
if not looping:
|
||||
num_frames = num_frames + num_refs * 4
|
||||
|
||||
vae.to(device)
|
||||
# Resize and rearrange the input image dimensions
|
||||
pose_latents = ref_latents = ref_latent = None
|
||||
if pose_images is not None:
|
||||
pose_images = pose_images[..., :3]
|
||||
if pose_images.shape[1] != H or pose_images.shape[2] != W:
|
||||
resized_pose_images = common_upscale(pose_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
|
||||
else:
|
||||
resized_pose_images = pose_images.permute(3, 0, 1, 2) # C, T, H, W
|
||||
resized_pose_images = resized_pose_images * 2 - 1
|
||||
pose_latents = vae.encode([resized_pose_images.to(device, vae.dtype)], device,tiled=tiled_vae)
|
||||
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]
|
||||
pad = torch.zeros(pose_latents.shape[0], pose_latents.shape[1], pad_len, pose_latents.shape[3], pose_latents.shape[4], device=pose_latents.device, dtype=pose_latents.dtype)
|
||||
pose_latents = torch.cat([pose_latents, pad], dim=2)
|
||||
print("pose_latents", pose_latents.shape)
|
||||
del resized_pose_images
|
||||
|
||||
bg_latents = None
|
||||
if bg_images is not None:
|
||||
if bg_images.shape[1] != H or bg_images.shape[2] != W:
|
||||
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]
|
||||
print("bg_latents", bg_latents.shape)
|
||||
del resized_bg_images
|
||||
else:
|
||||
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:
|
||||
resized_ref_images = common_upscale(ref_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
|
||||
else:
|
||||
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
|
||||
msk = torch.zeros(4, (frame_window_size - 1) // 4 + 1, lat_h, lat_w, device=device, dtype=vae.dtype)
|
||||
|
||||
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]
|
||||
del concatenated
|
||||
print("ref_latent", ref_latent.shape)
|
||||
|
||||
if mask is None:
|
||||
ref_mask = torch.zeros(1, num_frames, lat_h, lat_w, device=device, dtype=vae.dtype)
|
||||
else:
|
||||
ref_mask = 1 - mask[:num_frames]
|
||||
ref_mask = common_upscale(ref_mask.unsqueeze(1), lat_w, lat_h, "nearest", "disabled").squeeze(1)
|
||||
ref_mask = ref_mask.to(vae.dtype).to(device)
|
||||
ref_mask = ref_mask.unsqueeze(-1).permute(3, 0, 1, 2) # 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 not looping:
|
||||
if bg_images is not None:
|
||||
bg_latents_masked = torch.cat([ref_mask, 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)
|
||||
else:
|
||||
ref_latent = ref_latent_masked
|
||||
|
||||
if face_images is not None:
|
||||
face_images = face_images[..., :3]
|
||||
if face_images.shape[1] != 512 or face_images.shape[2] != 512:
|
||||
resized_face_images = common_upscale(face_images.movedim(-1, 1), 512, 512, "lanczos", "center").movedim(0, 1)
|
||||
else:
|
||||
resized_face_images = face_images.permute(3, 0, 1, 2) # B, C, T, H, W
|
||||
resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0)
|
||||
else:
|
||||
resized_face_images = torch.zeros(1, 3, num_frames, 512, 512, device=device, dtype=torch.float32)
|
||||
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])
|
||||
|
||||
if force_offload:
|
||||
vae.model.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
image_embeds = {
|
||||
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
|
||||
"negative_clip_context": clip_embeds.get("negative_clip_embeds", None) if clip_embeds is not None else None,
|
||||
"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_latent": ref_latent,
|
||||
"ref_image": resized_ref_images if ref_images is not None else None,
|
||||
"face_pixels": resized_face_images,
|
||||
"num_frames": num_frames,
|
||||
"target_shape": target_shape,
|
||||
"frame_window_size": frame_window_size,
|
||||
"lat_h": lat_h,
|
||||
"lat_w": lat_w,
|
||||
"vae": vae,
|
||||
"colormatch": colormatch,
|
||||
"looping": looping,
|
||||
"pose_strength": pose_strength,
|
||||
"face_strength": face_strength,
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
|
||||
class WanVideoEmptyEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -1813,6 +2030,9 @@ class WanVideoSampler:
|
||||
gguf_reader = model["gguf_reader"]
|
||||
control_lora = model["control_lora"]
|
||||
|
||||
vae = image_embeds.get("vae", None)
|
||||
tiled_vae = image_embeds.get("tiled_vae", False)
|
||||
|
||||
transformer_options = patcher.model_options.get("transformer_options", None)
|
||||
merge_loras = transformer_options["merge_loras"]
|
||||
|
||||
@@ -1990,6 +2210,7 @@ class WanVideoSampler:
|
||||
|
||||
drop_last = image_embeds.get("drop_last", False)
|
||||
has_ref = image_embeds.get("has_ref", False)
|
||||
|
||||
else: #t2v
|
||||
target_shape = image_embeds.get("target_shape", None)
|
||||
if target_shape is None:
|
||||
@@ -2105,11 +2326,12 @@ class WanVideoSampler:
|
||||
phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0)
|
||||
|
||||
|
||||
num_frames = image_embeds.get("num_frames", 0)
|
||||
#HuMo inputs
|
||||
humo_audio = image_embeds.get("humo_audio_emb", None)
|
||||
humo_audio_neg = image_embeds.get("humo_audio_emb_neg", None)
|
||||
humo_reference_count = image_embeds.get("humo_reference_count", 0)
|
||||
num_frames = image_embeds.get("num_frames", 0)
|
||||
|
||||
if humo_audio is not None:
|
||||
from .HuMo.nodes import get_audio_emb_window
|
||||
if not multitalk_sampling:
|
||||
@@ -2140,6 +2362,16 @@ class WanVideoSampler:
|
||||
if not isinstance(humo_audio_cfg_scale, list):
|
||||
humo_audio_cfg_scale = [humo_audio_cfg_scale] * (steps + 1)
|
||||
|
||||
# WanAnim inputs
|
||||
frame_window_size = image_embeds.get("frame_window_size", 77)
|
||||
wananimate_loop = image_embeds.get("looping", False)
|
||||
wananim_pose_latents = image_embeds.get("pose_latents", None)
|
||||
wananim_pose_strength = image_embeds.get("pose_strength", 1.0)
|
||||
wananim_face_strength = image_embeds.get("face_strength", 1.0)
|
||||
wananim_face_pixels = image_embeds.get("face_pixels", None)
|
||||
if image_cond is None:
|
||||
image_cond = image_embeds.get("ref_latent", None)
|
||||
|
||||
latent_video_length = noise.shape[1]
|
||||
|
||||
# Initialize FreeInit filter if enabled
|
||||
@@ -2349,7 +2581,7 @@ class WanVideoSampler:
|
||||
|
||||
# vid2vid
|
||||
noise_mask=original_image=None
|
||||
if samples is not None and not multitalk_sampling:
|
||||
if samples is not None and not multitalk_sampling and not wananimate_loop:
|
||||
saved_generator_state = samples.get("generator_state", None)
|
||||
if saved_generator_state is not None:
|
||||
seed_g.set_state(saved_generator_state)
|
||||
@@ -2470,38 +2702,7 @@ class WanVideoSampler:
|
||||
gc.collect()
|
||||
|
||||
#blockswap init
|
||||
if not transformer.patched_linear:
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name:
|
||||
param.data = param.data.to(device)
|
||||
if "control_adapter" in name:
|
||||
param.data = param.data.to(device)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
block_swap_args["offload_txt_emb"],
|
||||
block_swap_args["offload_img_emb"],
|
||||
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
|
||||
prefetch_blocks = block_swap_args.get("prefetch_blocks", 0),
|
||||
block_swap_debug = block_swap_args.get("block_swap_debug", False),
|
||||
)
|
||||
elif model["auto_cpu_offload"]:
|
||||
for module in transformer.modules():
|
||||
if hasattr(module, "offload"):
|
||||
module.offload()
|
||||
if hasattr(module, "onload"):
|
||||
module.onload()
|
||||
for block in transformer.blocks:
|
||||
block.modulation = torch.nn.Parameter(block.modulation.to(device))
|
||||
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
|
||||
else:
|
||||
transformer.to(device)
|
||||
init_blockswap(transformer, block_swap_args, model)
|
||||
|
||||
# Initialize Cache if enabled
|
||||
previous_cache_states = None
|
||||
@@ -2647,7 +2848,8 @@ class WanVideoSampler:
|
||||
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
|
||||
add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False,
|
||||
mtv_motion_tokens=None, s2v_audio_input=None, s2v_ref_motion=None, s2v_motion_frames=[1, 0], s2v_pose=None,
|
||||
humo_image_cond=None, humo_image_cond_neg=None, humo_audio=None, humo_audio_neg=None):
|
||||
humo_image_cond=None, humo_image_cond_neg=None, humo_audio=None, humo_audio_neg=None, wananim_pose_latents=None,
|
||||
wananim_face_pixels=None):
|
||||
nonlocal transformer
|
||||
z = z.to(dtype)
|
||||
autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear)
|
||||
@@ -2821,7 +3023,6 @@ class WanVideoSampler:
|
||||
humo_audio_input_neg = None
|
||||
else:
|
||||
humo_audio_input = humo_audio_input_neg = None
|
||||
|
||||
base_params = {
|
||||
'x': [z], # latent
|
||||
'y': [image_cond_input] if image_cond_input is not None else None, # image cond
|
||||
@@ -2865,8 +3066,10 @@ class WanVideoSampler:
|
||||
"s2v_audio_scale": s2v_audio_scale if s2v_audio_input is not None else 1.0, # speech-to-video audio scale
|
||||
"s2v_pose": s2v_pose if s2v_pose is not None else None, # speech-to-video pose control
|
||||
"s2v_motion_frames": s2v_motion_frames, # speech-to-video motion frames,
|
||||
"humo_audio": humo_audio_input, # humo audio input
|
||||
"humo_audio_scale": humo_audio_scale if humo_audio is not None else 1.0, # humo audio scale
|
||||
"humo_audio": humo_audio, # humo audio input
|
||||
"humo_audio_scale": humo_audio_scale if humo_audio is not None else 1,
|
||||
"wananim_pose_latents": wananim_pose_latents.to(device) if wananim_pose_latents is not None else None, # WanAnimate pose latents
|
||||
"wananim_face_pixel_values": wananim_face_pixels.to(device, torch.float32) if wananim_face_pixels is not None else None, # WanAnimate face images
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -3029,7 +3232,7 @@ class WanVideoSampler:
|
||||
from .latent_preview import prepare_callback #custom for tiny VAE previews
|
||||
callback = prepare_callback(patcher, len(timesteps))
|
||||
|
||||
if not multitalk_sampling and not framepack:
|
||||
if not multitalk_sampling and not framepack and not wananimate_loop:
|
||||
log.info(f"Input sequence length: {seq_len}")
|
||||
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
|
||||
|
||||
@@ -3117,7 +3320,7 @@ class WanVideoSampler:
|
||||
try:
|
||||
pbar = ProgressBar(len(timesteps))
|
||||
#region main loop start
|
||||
for idx, t in enumerate(tqdm(timesteps, disable=multitalk_sampling)):
|
||||
for idx, t in enumerate(tqdm(timesteps, disable=multitalk_sampling or wananimate_loop)):
|
||||
if flowedit_args is not None:
|
||||
if idx < skip_steps:
|
||||
continue
|
||||
@@ -3411,7 +3614,6 @@ class WanVideoSampler:
|
||||
|
||||
partial_s2v_audio_input = None
|
||||
if s2v_audio_input is not None:
|
||||
indices = (torch.arange(4 + 1) - 2) * 1
|
||||
audio_start = c[0] * 4
|
||||
audio_end = c[-1] * 4 + 1
|
||||
center_indices = torch.arange(audio_start, audio_end, 1)
|
||||
@@ -3425,6 +3627,20 @@ class WanVideoSampler:
|
||||
partial_add_cond = None
|
||||
if add_cond is not None:
|
||||
partial_add_cond = add_cond[:, :, c].to(device, dtype)
|
||||
|
||||
partial_wananim_face_pixels = partial_wananim_pose_latents = None
|
||||
if wananim_face_pixels is not None:
|
||||
start = c[0] * 4
|
||||
end = c[-1] * 4
|
||||
center_indices = torch.arange(start, end, 1)
|
||||
center_indices = torch.clamp(center_indices, min=0, max=wananim_face_pixels.shape[2] - 1)
|
||||
partial_wananim_face_pixels = wananim_face_pixels[:, :, center_indices].to(device, dtype)
|
||||
if wananim_pose_latents is not None:
|
||||
start = c[0]
|
||||
end = c[-1]
|
||||
center_indices = torch.arange(start, end, 1)
|
||||
center_indices = torch.clamp(center_indices, min=0, max=wananim_pose_latents.shape[2] - 1)
|
||||
partial_wananim_pose_latents = wananim_pose_latents[:, :, center_indices][:, :, :context_frames-1].to(device, dtype)
|
||||
|
||||
if len(timestep.shape) != 1:
|
||||
partial_timestep = timestep[:, c]
|
||||
@@ -3440,7 +3656,8 @@ class WanVideoSampler:
|
||||
partial_timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj,
|
||||
partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c, fantasy_portrait_input=partial_fantasy_portrait_input,
|
||||
mtv_motion_tokens=partial_mtv_motion_tokens, s2v_audio_input=partial_s2v_audio_input, s2v_motion_frames=[1, 0], s2v_pose=partial_s2v_pose,
|
||||
humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg,)
|
||||
humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg,
|
||||
wananim_face_pixels=partial_wananim_face_pixels, wananim_pose_latents=partial_wananim_pose_latents)
|
||||
|
||||
if cache_args is not None:
|
||||
self.window_tracker.cache_states[window_id] = new_teacache
|
||||
@@ -3461,7 +3678,7 @@ class WanVideoSampler:
|
||||
offloaded = False
|
||||
tiled_vae = image_embeds.get("tiled_vae", False)
|
||||
frame_num = clip_length = image_embeds.get("frame_window_size", 81)
|
||||
vae = image_embeds.get("vae", None)
|
||||
|
||||
clip_embeds = image_embeds.get("clip_context", None)
|
||||
if clip_embeds is not None:
|
||||
clip_embeds = clip_embeds.to(dtype)
|
||||
@@ -3680,37 +3897,8 @@ class WanVideoSampler:
|
||||
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
|
||||
elif gguf_reader is not None: #handle GGUF
|
||||
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
|
||||
|
||||
#blockswap init
|
||||
if not transformer.patched_linear:
|
||||
if block_swap_args is not None:
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name:
|
||||
param.data = param.data.to(device)
|
||||
if "control_adapter" in name:
|
||||
param.data = param.data.to(device)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
block_swap_args["offload_txt_emb"],
|
||||
block_swap_args["offload_img_emb"],
|
||||
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
|
||||
)
|
||||
elif model["auto_cpu_offload"]:
|
||||
for module in transformer.modules():
|
||||
if hasattr(module, "offload"):
|
||||
module.offload()
|
||||
if hasattr(module, "onload"):
|
||||
module.onload()
|
||||
for block in transformer.blocks:
|
||||
block.modulation = torch.nn.Parameter(block.modulation.to(device))
|
||||
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
|
||||
else:
|
||||
transformer.to(device)
|
||||
init_blockswap(transformer, block_swap_args, device, dtype)
|
||||
|
||||
# Use the appropriate prompt for this section
|
||||
if len(text_embeds["prompt_embeds"]) > 1:
|
||||
@@ -4072,6 +4260,296 @@ class WanVideoSampler:
|
||||
except:
|
||||
pass
|
||||
return {"video": gen_video_samples},
|
||||
# region wananimate loop
|
||||
elif wananimate_loop:
|
||||
# calculate frame counts
|
||||
total_frames = num_frames
|
||||
overlap = 0
|
||||
refert_num = 1
|
||||
|
||||
real_clip_len = frame_window_size - overlap
|
||||
last_clip_num = (total_frames - overlap) % real_clip_len
|
||||
extra = 0 if last_clip_num == 0 else real_clip_len - last_clip_num
|
||||
target_len = total_frames + extra
|
||||
target_latent_len = (target_len - 1) // 4 + 2
|
||||
latent_window_size = (frame_window_size - 1) // 4 + 1
|
||||
|
||||
from .utils import tensor_pingpong_pad
|
||||
|
||||
ref_latent = image_embeds.get("ref_latent", None)
|
||||
ref_images = image_embeds.get("ref_image", None)
|
||||
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
|
||||
#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}")
|
||||
if ref_masks is not None:
|
||||
ref_masks_in = tensor_pingpong_pad(ref_masks, target_latent_len)
|
||||
log.info(f"WanAnimate: Ref masks {ref_masks.shape} padded to shape {ref_masks.shape}")
|
||||
if bg_images is not None:
|
||||
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
|
||||
|
||||
colormatch = image_embeds.get("colormatch", "disabled")
|
||||
output_path = image_embeds.get("output_path", "")
|
||||
offload = image_embeds.get("force_offload", False)
|
||||
|
||||
lat_h, lat_w = noise.shape[2], noise.shape[3]
|
||||
start = start_latent = img_counter = step_iteration_count = iteration_count = 0
|
||||
end = frame_window_size
|
||||
end_latent = latent_window_size
|
||||
|
||||
estimated_iterations = target_len // frame_window_size
|
||||
callback = prepare_callback(patcher, estimated_iterations)
|
||||
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
|
||||
|
||||
# outer WanAnimate loop
|
||||
gen_video_list = []
|
||||
while True:
|
||||
if start >= total_frames:
|
||||
break
|
||||
|
||||
mm.soft_empty_cache()
|
||||
|
||||
mask_reft_len = 0 if start == 0 else refert_num
|
||||
|
||||
self.cache_state = [None, None]
|
||||
|
||||
if ref_latent is not None:
|
||||
vae.to(device)
|
||||
#ref_latents = vae.encode([ref_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0]
|
||||
#msk = torch.zeros(4, 1, lat_h, lat_w, device=device, dtype=dtype)
|
||||
#msk[:, :1] = 1
|
||||
#ref_latents = torch.cat([msk, ref_latents], dim=0) # 4+C 1 H W
|
||||
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:
|
||||
bg_image_slice = bg_images_in[:, start:end].to(device)
|
||||
else:
|
||||
bg_image_slice = torch.zeros(3, frame_window_size-mask_reft_len, lat_h * 8, lat_w * 8, device=device, dtype=vae.dtype)
|
||||
if mask_reft_len == 0:
|
||||
temporal_ref_latents = vae.encode([bg_image_slice], device,tiled=tiled_vae)[0]
|
||||
else:
|
||||
concatenated = torch.cat([current_ref_images.to(device, dtype=vae.dtype), bg_image_slice[:, mask_reft_len:]], dim=1)
|
||||
temporal_ref_latents = vae.encode([concatenated.to(device, vae.dtype)], device,tiled=tiled_vae)[0]
|
||||
msk[:, :mask_reft_len] = 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
|
||||
image_cond_in = torch.cat([ref_latent, temporal_ref_latents], dim=1) # 4+C T+trefs H W
|
||||
|
||||
noise = torch.randn(16, latent_window_size + 1, lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
|
||||
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
|
||||
|
||||
pose_input_slice = None
|
||||
if wananim_pose_latents is not None:
|
||||
pose_input_slice = wananim_pose_latents[:, :, start_latent:end_latent].to(device, dtype)
|
||||
# Pad if slice is too short
|
||||
if pose_input_slice.shape[2] < latent_window_size:
|
||||
log.info(f"WanAnimate: Padding pose latents from {pose_input_slice.shape} to length {latent_window_size}")
|
||||
pad_len = latent_window_size - pose_input_slice.shape[2]
|
||||
pad = torch.zeros(pose_input_slice.shape[0], pose_input_slice.shape[1], pad_len, pose_input_slice.shape[3], pose_input_slice.shape[4], device=pose_input_slice.device, dtype=pose_input_slice.dtype)
|
||||
pose_input_slice = torch.cat([pose_input_slice, pad], dim=2)
|
||||
pose_input_slice = pose_input_slice.to(device, dtype)
|
||||
|
||||
if samples is not None:
|
||||
input_samples = samples["samples"].squeeze(0).to(noise)
|
||||
# Check if we have enough frames in input_samples
|
||||
if latent_end_idx > input_samples.shape[1]:
|
||||
# We need more frames than available - pad the input_samples at the end
|
||||
pad_length = latent_end_idx - input_samples.shape[1]
|
||||
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
|
||||
input_samples = torch.cat([input_samples, last_frame], dim=1)
|
||||
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
|
||||
if noise_mask is not None:
|
||||
original_image = input_samples.to(device)
|
||||
|
||||
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
|
||||
|
||||
if add_noise_to_samples:
|
||||
latent_timestep = timesteps[0]
|
||||
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
|
||||
else:
|
||||
noise = input_samples
|
||||
|
||||
# diff diff prep
|
||||
noise_mask = samples.get("noise_mask", None)
|
||||
if noise_mask is not None:
|
||||
if len(noise_mask.shape) == 4:
|
||||
noise_mask = noise_mask.squeeze(1)
|
||||
if noise_mask.shape[0] < noise.shape[1]:
|
||||
noise_mask = noise_mask.repeat(noise.shape[1] // noise_mask.shape[0], 1, 1)
|
||||
else:
|
||||
noise_mask = noise_mask[latent_start_idx:latent_end_idx]
|
||||
noise_mask = torch.nn.functional.interpolate(
|
||||
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
|
||||
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
|
||||
mode='trilinear',
|
||||
align_corners=False
|
||||
).repeat(1, noise.shape[0], 1, 1, 1)
|
||||
|
||||
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
|
||||
thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device)
|
||||
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
|
||||
|
||||
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
|
||||
|
||||
# sample videos
|
||||
latent = noise
|
||||
|
||||
if offloaded:
|
||||
# Load weights
|
||||
if transformer.patched_linear and gguf_reader is None:
|
||||
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
|
||||
elif gguf_reader is not None: #handle GGUF
|
||||
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
|
||||
#blockswap init
|
||||
init_blockswap(transformer, block_swap_args, model)
|
||||
|
||||
# Use the appropriate prompt for this section
|
||||
if len(text_embeds["prompt_embeds"]) > 1:
|
||||
prompt_index = min(iteration_count, len(text_embeds["prompt_embeds"]) - 1)
|
||||
positive = [text_embeds["prompt_embeds"][prompt_index]]
|
||||
log.info(f"Using prompt index: {prompt_index}")
|
||||
else:
|
||||
positive = text_embeds["prompt_embeds"]
|
||||
|
||||
# uni3c slices
|
||||
if uni3c_embeds is not None:
|
||||
vae.to(device)
|
||||
# Pad original_images if needed
|
||||
num_frames = original_images.shape[2]
|
||||
required_frames = audio_end_idx - audio_start_idx
|
||||
if audio_end_idx > num_frames:
|
||||
pad_len = audio_end_idx - num_frames
|
||||
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
|
||||
padded_images = torch.cat([original_images, last_frame], dim=2)
|
||||
else:
|
||||
padded_images = original_images
|
||||
render_latent = vae.encode(
|
||||
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)
|
||||
pcd_data['render_latent'] = render_latent
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
# inner WanAnimate sampling loop
|
||||
sampling_pbar = tqdm(total=len(timesteps), desc=f"Frames {start}-{end}", position=0, leave=True)
|
||||
for i in range(len(timesteps)):
|
||||
timestep = timesteps[i]
|
||||
latent_model_input = latent.to(device)
|
||||
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
|
||||
timestep, i, cache_state=self.cache_state,
|
||||
image_cond = image_cond_in,
|
||||
wananim_face_pixels=face_images[:, :, start:end].to(device, torch.float32) if face_images is not None else None,
|
||||
wananim_pose_latents=pose_input_slice
|
||||
)
|
||||
if callback is not None:
|
||||
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
|
||||
callback(step_iteration_count, callback_latent, None, estimated_iterations*(len(timesteps)))
|
||||
del callback_latent
|
||||
|
||||
sampling_pbar.update(1)
|
||||
step_iteration_count += 1
|
||||
|
||||
latent = sample_scheduler.step(noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0).to(noise_pred.device), **scheduler_step_args)[0].squeeze(0)
|
||||
del noise_pred, latent_model_input, timestep
|
||||
|
||||
# differential diffusion inpaint
|
||||
if masks is not None:
|
||||
if i < len(timesteps) - 1:
|
||||
image_latent = add_noise(original_image.to(device), noise.to(device), timesteps[i+1])
|
||||
mask = masks[i].to(latent)
|
||||
latent = image_latent * mask + latent * (1-mask)
|
||||
|
||||
del noise
|
||||
if offload:
|
||||
offload_transformer(transformer)
|
||||
offloaded = True
|
||||
|
||||
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()
|
||||
vae.to(offload_device)
|
||||
|
||||
sampling_pbar.close()
|
||||
|
||||
# optional color correction
|
||||
if colormatch != "disabled":
|
||||
videos = videos.permute(1, 2, 3, 0).float().numpy()
|
||||
from color_matcher import ColorMatcher
|
||||
cm = ColorMatcher()
|
||||
cm_result_list = []
|
||||
for img in videos:
|
||||
cm_result = cm.transfer(src=img, ref=ref_images.permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
|
||||
cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype))
|
||||
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
|
||||
|
||||
current_ref_images = videos[:, -1:].clone().detach()
|
||||
|
||||
# optionally save generated samples to disk
|
||||
if output_path:
|
||||
video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8')
|
||||
num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num
|
||||
log.info(f"Saving {num_frames_to_save} generated frames to {output_path}")
|
||||
start_idx = 0 if is_first_clip else cur_motion_frames_num
|
||||
for i in range(start_idx, video_np.shape[0]):
|
||||
im = Image.fromarray(video_np[i])
|
||||
im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png"))
|
||||
img_counter += 1
|
||||
else:
|
||||
gen_video_list.append(videos)
|
||||
|
||||
del videos
|
||||
|
||||
iteration_count += 1
|
||||
start += frame_window_size
|
||||
end += frame_window_size
|
||||
start_latent += latent_window_size
|
||||
end_latent += latent_window_size
|
||||
|
||||
if not output_path:
|
||||
gen_video_samples = torch.cat(gen_video_list, dim=1)
|
||||
else:
|
||||
gen_video_samples = torch.zeros(3, 1, 64, 64) # dummy output
|
||||
|
||||
if force_offload:
|
||||
if not model["auto_cpu_offload"]:
|
||||
offload_transformer(transformer)
|
||||
try:
|
||||
print_memory(device)
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path},
|
||||
|
||||
#region normal inference
|
||||
else:
|
||||
@@ -4081,7 +4559,9 @@ class WanVideoSampler:
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, multitalk_audio_embeds=multitalk_audio_embeds, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input,
|
||||
humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg)
|
||||
humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg,
|
||||
wananim_face_pixels=wananim_face_pixels, wananim_pose_latents=wananim_pose_latents,
|
||||
)
|
||||
if bidirectional_sampling:
|
||||
noise_pred_flipped, self.cache_state = predict_with_cfg(
|
||||
latent_model_input_flipped,
|
||||
@@ -4188,6 +4668,8 @@ class WanVideoSampler:
|
||||
latent = latent[:,:-phantom_latents.shape[1]]
|
||||
if humo_reference_count > 0:
|
||||
latent = latent[:,:-humo_reference_count]
|
||||
if wananim_pose_latents is not None:
|
||||
latent = latent[:, 1:]
|
||||
|
||||
cache_states = None
|
||||
if cache_args is not None:
|
||||
@@ -4469,7 +4951,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddControlEmbeds": WanVideoAddControlEmbeds,
|
||||
"WanVideoAddMTVMotion": WanVideoAddMTVMotion,
|
||||
"WanVideoRoPEFunction": WanVideoRoPEFunction,
|
||||
"WanVideoAddPusaNoise": WanVideoAddPusaNoise
|
||||
"WanVideoAddPusaNoise": WanVideoAddPusaNoise,
|
||||
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoSampler": "WanVideo Sampler",
|
||||
@@ -4506,4 +4989,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddMTVMotion": "WanVideo MTV Crafter Motion",
|
||||
"WanVideoRoPEFunction": "WanVideo RoPE Function",
|
||||
"WanVideoAddPusaNoise": "WanVideo Add Pusa Noise",
|
||||
"WanVideoAnimateEmbeds": "WanVideo Animate Embeds",
|
||||
}
|
||||
|
||||
@@ -761,7 +761,7 @@ class WanVideoSetLoRAs:
|
||||
def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None):
|
||||
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding",
|
||||
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob"}
|
||||
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob", "motion_encoder"}
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
pbar = ProgressBar(param_count)
|
||||
cnt = 0
|
||||
@@ -851,7 +851,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
dtype_to_use = weight_dtype if sd[name.replace("_orig_mod.", "")].dtype == weight_dtype else dtype_to_use
|
||||
if "modulation" in name or "norm" in name or "bias" in name or "img_emb" in name:
|
||||
dtype_to_use = base_dtype
|
||||
if "patch_embedding" in name:
|
||||
if "patch_embedding" in name or "motion_encoder" in name or "face_encoder" in name:
|
||||
dtype_to_use = torch.float32
|
||||
|
||||
load_device = transformer_load_device
|
||||
@@ -1116,6 +1116,7 @@ class WanVideoModelLoader:
|
||||
ffn2_dim = sd["blocks.0.ffn.2.weight"].shape[1]
|
||||
|
||||
is_humo = "audio_proj.audio_proj_glob_1.layer.weight" in sd
|
||||
is_wananimate = "pose_patch_embedding.weight" in sd
|
||||
|
||||
model_type = "t2v"
|
||||
if "audio_injector.injector.0.k.weight" in sd:
|
||||
@@ -1219,6 +1220,7 @@ class WanVideoModelLoader:
|
||||
"rope_func": "comfy",
|
||||
"main_device": device,
|
||||
"offload_device": offload_device,
|
||||
"dtype": base_dtype,
|
||||
"teacache_coefficients": teacache_coefficients_map[model_variant],
|
||||
"magcache_ratios": magcache_ratios_map[model_variant],
|
||||
"vace_layers": vace_layers,
|
||||
@@ -1232,6 +1234,7 @@ class WanVideoModelLoader:
|
||||
"cond_dim": sd["cond_encoder.weight"].shape[1] if "cond_encoder.weight" in sd else 0,
|
||||
"zero_timestep": model_type == "s2v",
|
||||
"humo_audio": is_humo,
|
||||
"is_wananimate": is_wananimate,
|
||||
|
||||
}
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ import torch
|
||||
import numpy as np
|
||||
from comfy.utils import common_upscale
|
||||
from .utils import log
|
||||
from einops import rearrange
|
||||
|
||||
try:
|
||||
from server import PromptServer
|
||||
@@ -476,6 +477,73 @@ class WanVideoPassImagesFromSamples:
|
||||
video.clamp_(-1.0, 1.0)
|
||||
video.add_(1.0).div_(2.0)
|
||||
return video.cpu().float(), samples.get("output_path", "")
|
||||
|
||||
|
||||
class FaceMaskFromPoseKeypoints:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
input_types = {
|
||||
"required": {
|
||||
"pose_kps": ("POSE_KEYPOINT",),
|
||||
"person_index": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1, "tooltip": "Index of the person to start with"}),
|
||||
}
|
||||
}
|
||||
return input_types
|
||||
RETURN_TYPES = ("MASK",)
|
||||
FUNCTION = "createmask"
|
||||
CATEGORY = "ControlNet Preprocessors/Pose Keypoint Postprocess"
|
||||
|
||||
def createmask(self, pose_kps, person_index):
|
||||
pose_frames = pose_kps
|
||||
prev_center = None
|
||||
np_frames = []
|
||||
for i, pose_frame in enumerate(pose_frames):
|
||||
selected_idx, prev_center = self.select_closest_person(pose_frame, person_index if i == 0 else prev_center)
|
||||
np_frames.append(self.draw_kps(pose_frame, selected_idx))
|
||||
np_frames = np.stack(np_frames, axis=0)
|
||||
tensor = torch.from_numpy(np_frames).float() / 255.
|
||||
print("tensor.shape:", tensor.shape)
|
||||
tensor = tensor[:, :, :, 0]
|
||||
return (tensor,)
|
||||
|
||||
def select_closest_person(self, pose_frame, prev_center_or_index):
|
||||
people = pose_frame["people"]
|
||||
if not people:
|
||||
return -1, None
|
||||
centers = []
|
||||
for person in people:
|
||||
kps = np.array(person["face_keypoints_2d"])
|
||||
n = len(kps) // 3
|
||||
facial_kps = rearrange(kps, "(n c) -> n c", n=n, c=3)[:, :2]
|
||||
center = facial_kps.mean(axis=0)
|
||||
centers.append(center)
|
||||
if isinstance(prev_center_or_index, (int, np.integer)):
|
||||
# First frame: use person_index
|
||||
idx = prev_center_or_index if 0 <= prev_center_or_index < len(people) else 0
|
||||
return idx, centers[idx]
|
||||
else:
|
||||
# Find closest to previous center
|
||||
prev_center = np.array(prev_center_or_index)
|
||||
dists = [np.linalg.norm(center - prev_center) for center in centers]
|
||||
idx = int(np.argmin(dists))
|
||||
return idx, centers[idx]
|
||||
|
||||
def draw_kps(self, pose_frame, person_index):
|
||||
import cv2
|
||||
width, height = pose_frame["canvas_width"], pose_frame["canvas_height"]
|
||||
canvas = np.zeros((height, width, 3), dtype=np.uint8)
|
||||
people = pose_frame["people"]
|
||||
if person_index < 0 or person_index >= len(people):
|
||||
return canvas # Out of bounds, return blank
|
||||
person = people[person_index]
|
||||
n = len(person["face_keypoints_2d"]) // 3
|
||||
facial_kps = rearrange(np.array(person["face_keypoints_2d"]), "(n c) -> n c", n=n, c=3)[:, :2]
|
||||
facial_kps = facial_kps.astype(np.int32)
|
||||
part_color = (255, 255, 255)
|
||||
|
||||
outer_contour = facial_kps[:17]
|
||||
cv2.fillPoly(canvas, pts=[outer_contour], color=part_color)
|
||||
return canvas
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
|
||||
@@ -488,6 +556,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoSigmaToStep": WanVideoSigmaToStep,
|
||||
"NormalizeAudioLoudness": NormalizeAudioLoudness,
|
||||
"WanVideoPassImagesFromSamples": WanVideoPassImagesFromSamples,
|
||||
"FaceMaskFromPoseKeypoints": FaceMaskFromPoseKeypoints,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
|
||||
@@ -500,4 +569,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoSigmaToStep": "WanVideo Sigma To Step",
|
||||
"NormalizeAudioLoudness": "Normalize Audio Loudness",
|
||||
"WanVideoPassImagesFromSamples": "WanVideo Pass Images From Samples",
|
||||
"FaceMaskFromPoseKeypoints": "Face Mask From Pose Keypoints",
|
||||
}
|
||||
@@ -3,6 +3,7 @@ import torch
|
||||
import logging
|
||||
import math
|
||||
from tqdm import tqdm
|
||||
from copy import deepcopy
|
||||
import types, collections
|
||||
from comfy.utils import ProgressBar, copy_to_param, set_attr_param
|
||||
from comfy.model_patcher import get_key_weight, string_to_seed
|
||||
@@ -538,4 +539,32 @@ def get_raag_guidance(noise_pred_cond, noise_pred_uncond, w_max, alpha=1.0, eps=
|
||||
ratio = norm_delta / (norm_uncond + eps)
|
||||
ratio_mean = ratio.mean().item()
|
||||
adaptive_w = 1.0 + (w_max - 1.0) * math.exp(-alpha * ratio_mean)
|
||||
return adaptive_w
|
||||
return adaptive_w
|
||||
|
||||
def tensor_pingpong_pad(video, target_len):
|
||||
"""
|
||||
Pads a video tensor along the frame dimension (dim=2) in a ping-pong fashion.
|
||||
video: torch.Tensor of shape [B, C, F, H, W]
|
||||
target_len: desired number of frames
|
||||
Returns: padded tensor of shape [B, C, target_len, H, W]
|
||||
"""
|
||||
in_dims = len(video.shape)
|
||||
if in_dims == 4:
|
||||
video = video.unsqueeze(0)
|
||||
B, C, F, H, W = video.shape
|
||||
idx = 0
|
||||
flip = False
|
||||
indices = []
|
||||
while len(indices) < target_len:
|
||||
indices.append(idx)
|
||||
if flip:
|
||||
idx -= 1
|
||||
else:
|
||||
idx += 1
|
||||
if idx == 0 or idx == F - 1:
|
||||
flip = not flip
|
||||
indices = indices[:target_len]
|
||||
padded_video = video[:, :, indices, :, :]
|
||||
if in_dims == 4:
|
||||
padded_video = padded_video.squeeze(0)
|
||||
return padded_video
|
||||
@@ -584,7 +584,7 @@ class LoRALinearLayer(nn.Module):
|
||||
down_hidden_states = self.down(hidden_states.to(dtype))
|
||||
up_hidden_states = self.up(down_hidden_states) * self.strength
|
||||
return up_hidden_states.to(orig_dtype)
|
||||
|
||||
|
||||
#region crossattn
|
||||
class WanT2VCrossAttention(WanSelfAttention):
|
||||
|
||||
@@ -1445,6 +1445,7 @@ class WanModel(torch.nn.Module):
|
||||
rope_func='comfy',
|
||||
main_device=torch.device('cuda'),
|
||||
offload_device=torch.device('cpu'),
|
||||
dtype=torch.float16,
|
||||
teacache_coefficients=[],
|
||||
magcache_ratios=[],
|
||||
vace_layers=None,
|
||||
@@ -1464,6 +1465,9 @@ class WanModel(torch.nn.Module):
|
||||
audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39],
|
||||
zero_timestep=False,
|
||||
humo_audio=False,
|
||||
# WanAnimate
|
||||
is_wananimate=False,
|
||||
motion_encoder_dim=512,
|
||||
):
|
||||
r"""
|
||||
Initialize the diffusion model backbone.
|
||||
@@ -1575,6 +1579,8 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
self.humo_audio = humo_audio
|
||||
|
||||
self.motion_encoder_dim = motion_encoder_dim
|
||||
|
||||
# embeddings
|
||||
self.patch_embedding = nn.Conv3d(
|
||||
in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
@@ -1714,7 +1720,23 @@ class WanModel(torch.nn.Module):
|
||||
from ...HuMo.audio_proj import AudioProjModel
|
||||
self.audio_proj = AudioProjModel(seq_len=8, blocks=5, channels=1280,
|
||||
intermediate_dim=512, output_dim=1536, context_tokens=16)
|
||||
|
||||
# WanAnimate
|
||||
self.motion_encoder = self.pose_patch_embedding = self.face_encoder = self.face_adapter = None
|
||||
if is_wananimate:
|
||||
from .wananimate.motion_encoder import MotionExtractor
|
||||
from .wananimate.face_blocks import FaceEncoder, FaceAdapter
|
||||
self.pose_patch_embedding = nn.Conv3d(16, dim, kernel_size=patch_size, stride=patch_size)
|
||||
self.motion_encoder = MotionExtractor()
|
||||
self.face_adapter = FaceAdapter(
|
||||
num_heads=self.num_heads,
|
||||
feature_dim=self.dim,
|
||||
num_adapter_layers=self.num_layers // 5,
|
||||
)
|
||||
self.face_encoder = FaceEncoder(
|
||||
in_dim=motion_encoder_dim,
|
||||
out_dim=self.dim,
|
||||
num_heads=4,
|
||||
)
|
||||
|
||||
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None, prefetch_blocks=0, block_swap_debug=False):
|
||||
# Clamp blocks_to_swap to valid range
|
||||
@@ -1850,6 +1872,44 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
return x
|
||||
|
||||
def wananimate_pose_embedding(self, x, pose_latents, strength=1.0):
|
||||
pose_latents = [self.pose_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in pose_latents]
|
||||
for x_, pose_latents_ in zip(x, pose_latents):
|
||||
x_[:, :, 1:].add_(pose_latents_, alpha=strength)
|
||||
return x
|
||||
|
||||
|
||||
def wananimate_face_embedding(self, face_pixel_values):
|
||||
b,c,T,h,w = face_pixel_values.shape
|
||||
face_pixel_values = rearrange(face_pixel_values, "b c t h w -> (b t) c h w")
|
||||
|
||||
encode_bs = 8
|
||||
face_pixel_values_tmp = []
|
||||
self.motion_encoder.to(self.main_device)
|
||||
for i in range(math.ceil(face_pixel_values.shape[0]/encode_bs)):
|
||||
face_pixel_values_tmp.append(self.motion_encoder(face_pixel_values[i*encode_bs:(i+1)*encode_bs]))
|
||||
del face_pixel_values
|
||||
self.motion_encoder.to(self.offload_device)
|
||||
|
||||
motion_vec = rearrange(torch.cat(face_pixel_values_tmp), "(b t) c -> b t c", t=T)
|
||||
del face_pixel_values_tmp
|
||||
self.face_encoder.to(self.main_device)
|
||||
motion_vec = self.face_encoder(motion_vec)
|
||||
self.face_encoder.to(self.offload_device)
|
||||
|
||||
B, L, H, C = motion_vec.shape
|
||||
pad_face = torch.zeros(B, 1, H, C, device=motion_vec.device, dtype=motion_vec.dtype)
|
||||
return torch.cat([pad_face, motion_vec], dim=1)
|
||||
|
||||
|
||||
def wananimate_forward(self, block_idx, x, motion_vec, strength=1.0, motion_masks=None):
|
||||
if block_idx % 5 == 0:
|
||||
adapter_args = [x, motion_vec, motion_masks]
|
||||
residual_out = self.face_adapter.fuser_blocks[block_idx // 5](*adapter_args)
|
||||
return x.add(residual_out, alpha=strength)
|
||||
return x
|
||||
|
||||
|
||||
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond_shape=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
|
||||
patch_size = self.patch_size
|
||||
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
|
||||
@@ -1949,6 +2009,10 @@ class WanModel(torch.nn.Module):
|
||||
s2v_motion_frames=[1, 0],
|
||||
humo_audio=None,
|
||||
humo_audio_scale=1.0,
|
||||
wananim_pose_latents=None,
|
||||
wananim_face_pixel_values=None,
|
||||
wananim_pose_strength=1.0,
|
||||
wananim_face_strength=1.0,
|
||||
|
||||
):
|
||||
r"""
|
||||
@@ -2045,10 +2109,20 @@ class WanModel(torch.nn.Module):
|
||||
self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype)
|
||||
for u in x
|
||||
]
|
||||
|
||||
|
||||
# WanAnimate
|
||||
motion_vec = None
|
||||
if wananim_face_pixel_values is not None:
|
||||
motion_vec = self.wananimate_face_embedding(wananim_face_pixel_values).to(x[0].dtype)
|
||||
|
||||
if wananim_pose_latents is not None:
|
||||
x = self.wananimate_pose_embedding(x, wananim_pose_latents, strength=wananim_pose_strength)
|
||||
|
||||
# s2v pose embedding
|
||||
if s2v_pose is not None:
|
||||
x[0] = x[0] + self.cond_encoder(s2v_pose.to(self.cond_encoder.weight.dtype)).to(x[0].dtype)
|
||||
|
||||
# Fun camera
|
||||
if self.control_adapter is not None and fun_camera is not None:
|
||||
fun_camera = self.control_adapter(fun_camera)
|
||||
x = [u + v for u, v in zip(x, fun_camera)]
|
||||
@@ -2550,6 +2624,8 @@ class WanModel(torch.nn.Module):
|
||||
x, x_ip = block(x, x_ip=x_ip, **kwargs) #run block
|
||||
if self.audio_injector is not None and s2v_audio_input is not None:
|
||||
x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v
|
||||
if self.motion_encoder is not None and motion_vec is not None:
|
||||
x = self.wananimate_forward(b, x, motion_vec, strength=wananim_face_strength)
|
||||
if self.block_swap_debug:
|
||||
compute_end = time.perf_counter()
|
||||
compute_time = compute_end - compute_start
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
from torch import nn
|
||||
import torch
|
||||
from einops import rearrange
|
||||
import torch.nn.functional as F
|
||||
from ..attention import attention
|
||||
|
||||
class CausalConv1d(nn.Module):
|
||||
def __init__(self, chan_in, chan_out, kernel_size=3, stride=1, dilation=1, pad_mode="replicate", **kwargs):
|
||||
super().__init__()
|
||||
|
||||
self.pad_mode = pad_mode
|
||||
padding = (kernel_size - 1, 0) # T
|
||||
self.time_causal_padding = padding
|
||||
|
||||
self.conv = nn.Conv1d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class FaceEncoder(nn.Module):
|
||||
def __init__(self, in_dim: int, out_dim: int, num_heads: int, dtype=None, device=None):
|
||||
super().__init__()
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.conv1_local = CausalConv1d(in_dim, 1024 * num_heads, 3, stride=1)
|
||||
self.norm1 = nn.LayerNorm(1024, elementwise_affine=False, eps=1e-6, device=device, dtype=dtype)
|
||||
self.act = nn.SiLU()
|
||||
self.conv2 = CausalConv1d(1024, 1024, 3, stride=2)
|
||||
self.conv3 = CausalConv1d(1024, 1024, 3, stride=2)
|
||||
|
||||
self.out_proj = nn.Linear(1024, out_dim)
|
||||
|
||||
self.norm2 = nn.LayerNorm(1024, elementwise_affine=False, eps=1e-6, device=device, dtype=dtype)
|
||||
self.norm3 = nn.LayerNorm(1024, elementwise_affine=False, eps=1e-6, device=device, dtype=dtype)
|
||||
|
||||
self.padding_tokens = nn.Parameter(torch.zeros(1, 1, 1, out_dim))
|
||||
|
||||
def forward(self, x):
|
||||
x = rearrange(x, "b t c -> b c t")
|
||||
b = x.shape[0]
|
||||
|
||||
x = self.conv1_local(x)
|
||||
x = rearrange(x, "b (n c) t -> (b n) t c", n=self.num_heads)
|
||||
|
||||
x = self.norm1(x)
|
||||
x = self.act(x)
|
||||
x = rearrange(x, "b t c -> b c t")
|
||||
x = self.conv2(x)
|
||||
x = rearrange(x, "b c t -> b t c")
|
||||
x = self.norm2(x)
|
||||
x = self.act(x)
|
||||
x = rearrange(x, "b t c -> b c t")
|
||||
x = self.conv3(x)
|
||||
x = rearrange(x, "b c t -> b t c")
|
||||
x = self.norm3(x)
|
||||
x = self.act(x)
|
||||
x = self.out_proj(x)
|
||||
x = rearrange(x, "(b n) t c -> b t n c", b=b)
|
||||
padding = self.padding_tokens.repeat(b, x.shape[1], 1, 1)
|
||||
|
||||
return torch.cat([x, padding], dim=-2)
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim, elementwise_affine=True, eps=1e-6, device=None, dtype=None):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
if elementwise_affine:
|
||||
self.weight = nn.Parameter(torch.ones(dim, device=device, dtype=dtype))
|
||||
|
||||
def _norm(self, x):
|
||||
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
||||
|
||||
def forward(self, x):
|
||||
output = self._norm(x.float()).type_as(x)
|
||||
if hasattr(self, "weight"):
|
||||
output = output * self.weight
|
||||
return output
|
||||
|
||||
|
||||
class FaceAdapter(nn.Module):
|
||||
def __init__(self, feature_dim, num_heads, num_adapter_layers=1, dtype=None, device=None):
|
||||
super().__init__()
|
||||
self.fuser_blocks = nn.ModuleList([FaceBlock(feature_dim, num_heads, device=device, dtype=dtype) for _ in range(num_adapter_layers)])
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
motion_embed: torch.Tensor,
|
||||
idx: int,
|
||||
) -> torch.Tensor:
|
||||
|
||||
return self.fuser_blocks[idx](x, motion_embed)
|
||||
|
||||
|
||||
class FaceBlock(nn.Module):
|
||||
def __init__(self, feature_dim, num_heads, dtype=None, device=None):
|
||||
super().__init__()
|
||||
|
||||
self.feature_dim = feature_dim
|
||||
self.num_heads = num_heads
|
||||
head_dim = feature_dim // num_heads
|
||||
|
||||
self.linear1_kv = nn.Linear(feature_dim, feature_dim * 2, device=device, dtype=dtype)
|
||||
self.linear1_q = nn.Linear(feature_dim, feature_dim, device=device, dtype=dtype)
|
||||
self.linear2 = nn.Linear(feature_dim, feature_dim, device=device, dtype=dtype)
|
||||
|
||||
self.q_norm = (RMSNorm(head_dim, elementwise_affine=True, eps=1e-6, device=device, dtype=dtype))
|
||||
self.k_norm = (RMSNorm(head_dim, elementwise_affine=True, eps=1e-6, device=device, dtype=dtype))
|
||||
|
||||
self.pre_norm_feat = nn.LayerNorm(feature_dim, elementwise_affine=False, eps=1e-6, device=device, dtype=dtype)
|
||||
self.pre_norm_motion = nn.LayerNorm(feature_dim, elementwise_affine=False, eps=1e-6, device=device, dtype=dtype)
|
||||
|
||||
|
||||
def forward(self, x, motion_vec, motion_mask=None):
|
||||
B, T, N, C = motion_vec.shape
|
||||
|
||||
x_motion = self.pre_norm_motion(motion_vec)
|
||||
x_feat = self.pre_norm_feat(x)
|
||||
|
||||
kv = self.linear1_kv(x_motion)
|
||||
q = self.linear1_q(x_feat)
|
||||
|
||||
k, v = rearrange(kv, "B L N (K H D) -> K B L N H D", K=2, H=self.num_heads)
|
||||
q = rearrange(q, "B S (H D) -> B S H D", H=self.num_heads)
|
||||
|
||||
q = self.q_norm(q).to(v)
|
||||
k = self.k_norm(k).to(v)
|
||||
|
||||
k = rearrange(k, "B L N H D -> (B L) N H D")
|
||||
v = rearrange(v, "B L N H D -> (B L) N H D")
|
||||
q = rearrange(q, "B (L S) H D -> (B L) S H D", L=T)
|
||||
|
||||
attn = attention(q, k, v)
|
||||
attn = attn.reshape(attn.shape[0], attn.shape[1], -1)
|
||||
attn = rearrange(attn, "(B L) S C -> B (L S) C", L=T)
|
||||
output = self.linear2(attn)
|
||||
|
||||
if motion_mask is not None:
|
||||
output = output * rearrange(motion_mask, "B T H W -> B (T H W)").unsqueeze(-1)
|
||||
|
||||
return output
|
||||
@@ -0,0 +1,176 @@
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
import math
|
||||
|
||||
# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/ops/upfirdn2d/upfirdn2d.py#L162
|
||||
def upfirdn2d_native(input, kernel, up_x, up_y, down_x, down_y, pad_x0, pad_x1, pad_y0, pad_y1):
|
||||
_, minor, in_h, in_w = input.shape
|
||||
kernel_h, kernel_w = kernel.shape
|
||||
|
||||
out = input.view(-1, minor, in_h, 1, in_w, 1)
|
||||
out = F.pad(out, [0, up_x - 1, 0, 0, 0, up_y - 1, 0, 0])
|
||||
out = out.view(-1, minor, in_h * up_y, in_w * up_x)
|
||||
|
||||
out = F.pad(out, [max(pad_x0, 0), max(pad_x1, 0), max(pad_y0, 0), max(pad_y1, 0)])
|
||||
out = out[:, :, max(-pad_y0, 0): out.shape[2] - max(-pad_y1, 0), max(-pad_x0, 0): out.shape[3] - max(-pad_x1, 0)]
|
||||
|
||||
out = out.reshape([-1, 1, in_h * up_y + pad_y0 + pad_y1, in_w * up_x + pad_x0 + pad_x1])
|
||||
w = torch.flip(kernel, [0, 1]).view(1, 1, kernel_h, kernel_w)
|
||||
out = F.conv2d(out, w)
|
||||
out = out.reshape(-1, minor, in_h * up_y + pad_y0 + pad_y1 - kernel_h + 1, in_w * up_x + pad_x0 + pad_x1 - kernel_w + 1)
|
||||
return out[:, :, ::down_y, ::down_x]
|
||||
|
||||
def upfirdn2d(input, kernel, up=1, down=1, pad=(0, 0)):
|
||||
return upfirdn2d_native(input, kernel, up, up, down, down, pad[0], pad[1], pad[0], pad[1])
|
||||
|
||||
# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/ops/fused_act/fused_act.py#L81
|
||||
class FusedLeakyReLU(torch.nn.Module):
|
||||
def __init__(self, channel, negative_slope=0.2, scale=2 ** 0.5):
|
||||
super().__init__()
|
||||
self.bias = torch.nn.Parameter(torch.zeros(1, channel, 1, 1))
|
||||
self.negative_slope = negative_slope
|
||||
self.scale = scale
|
||||
|
||||
def forward(self, input):
|
||||
return fused_leaky_relu(input, self.bias, self.negative_slope, self.scale)
|
||||
|
||||
def fused_leaky_relu(input, bias, negative_slope=0.2, scale=2 ** 0.5):
|
||||
return F.leaky_relu(input + bias, negative_slope) * scale
|
||||
|
||||
class Blur(torch.nn.Module):
|
||||
def __init__(self, kernel, pad):
|
||||
super().__init__()
|
||||
kernel = torch.tensor(kernel, dtype=torch.float32)
|
||||
kernel = kernel[None, :] * kernel[:, None]
|
||||
kernel = kernel / kernel.sum()
|
||||
self.register_buffer('kernel', kernel)
|
||||
self.pad = pad
|
||||
|
||||
def forward(self, input):
|
||||
return upfirdn2d(input, self.kernel, pad=self.pad)
|
||||
|
||||
#https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/archs/stylegan2_arch.py#L590
|
||||
class ScaledLeakyReLU(torch.nn.Module):
|
||||
def __init__(self, negative_slope=0.2):
|
||||
super().__init__()
|
||||
self.negative_slope = negative_slope
|
||||
|
||||
def forward(self, input):
|
||||
return F.leaky_relu(input, negative_slope=self.negative_slope)
|
||||
|
||||
# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/archs/stylegan2_arch.py#L605
|
||||
class EqualConv2d(torch.nn.Module):
|
||||
def __init__(self, in_channel, out_channel, kernel_size, stride=1, padding=0, bias=True):
|
||||
super().__init__()
|
||||
self.weight = torch.nn.Parameter(torch.randn(out_channel, in_channel, kernel_size, kernel_size))
|
||||
self.scale = 1 / math.sqrt(in_channel * kernel_size ** 2)
|
||||
self.stride = stride
|
||||
self.padding = padding
|
||||
self.bias = torch.nn.Parameter(torch.zeros(out_channel)) if bias else None
|
||||
|
||||
def forward(self, input):
|
||||
return F.conv2d(input, self.weight * self.scale, bias=self.bias, stride=self.stride, padding=self.padding)
|
||||
|
||||
# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/archs/stylegan2_arch.py#L134
|
||||
class EqualLinear(torch.nn.Module):
|
||||
def __init__(self, in_dim, out_dim, bias=True, bias_init=0, lr_mul=1, activation=None):
|
||||
super().__init__()
|
||||
self.weight = torch.nn.Parameter(torch.randn(out_dim, in_dim).div_(lr_mul))
|
||||
self.bias = torch.nn.Parameter(torch.zeros(out_dim).fill_(bias_init)) if bias else None
|
||||
self.activation = activation
|
||||
self.scale = (1 / math.sqrt(in_dim)) * lr_mul
|
||||
self.lr_mul = lr_mul
|
||||
|
||||
def forward(self, input):
|
||||
if self.activation:
|
||||
out = F.linear(input, self.weight * self.scale)
|
||||
return fused_leaky_relu(out, self.bias * self.lr_mul)
|
||||
return F.linear(input, self.weight * self.scale, bias=self.bias * self.lr_mul)
|
||||
|
||||
# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/archs/stylegan2_arch.py#L654
|
||||
class ConvLayer(torch.nn.Sequential):
|
||||
def __init__(self, in_channel, out_channel, kernel_size, downsample=False, blur_kernel=[1, 3, 3, 1], bias=True, activate=True):
|
||||
layers = []
|
||||
|
||||
if downsample:
|
||||
factor = 2
|
||||
p = (len(blur_kernel) - factor) + (kernel_size - 1)
|
||||
layers.append(Blur(blur_kernel, pad=((p + 1) // 2, p // 2)))
|
||||
stride, padding = 2, 0
|
||||
else:
|
||||
stride, padding = 1, kernel_size // 2
|
||||
|
||||
layers.append(EqualConv2d(in_channel, out_channel, kernel_size, padding=padding, stride=stride, bias=bias and not activate))
|
||||
|
||||
if activate:
|
||||
layers.append(FusedLeakyReLU(out_channel) if bias else ScaledLeakyReLU(0.2))
|
||||
|
||||
super().__init__(*layers)
|
||||
|
||||
# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/archs/stylegan2_arch.py#L704
|
||||
class ResBlock(torch.nn.Module):
|
||||
def __init__(self, in_channel, out_channel):
|
||||
super().__init__()
|
||||
self.conv1 = ConvLayer(in_channel, in_channel, 3)
|
||||
self.conv2 = ConvLayer(in_channel, out_channel, 3, downsample=True)
|
||||
self.skip = ConvLayer(in_channel, out_channel, 1, downsample=True, activate=False, bias=False)
|
||||
|
||||
def forward(self, input):
|
||||
out = self.conv2(self.conv1(input))
|
||||
skip = self.skip(input)
|
||||
return (out + skip) / math.sqrt(2)
|
||||
|
||||
|
||||
class AppearanceEncoder(torch.nn.Module):
|
||||
def __init__(self, w_dim=512):
|
||||
super().__init__()
|
||||
|
||||
self.convs = torch.nn.ModuleList([
|
||||
ConvLayer(3, 32, 1), ResBlock(32, 64),
|
||||
ResBlock(64, 128), ResBlock(128, 256),
|
||||
ResBlock(256, 512), ResBlock(512, 512),
|
||||
ResBlock(512, 512), ResBlock(512, 512),
|
||||
EqualConv2d(512, w_dim, 4, padding=0, bias=False)
|
||||
])
|
||||
|
||||
def forward(self, x):
|
||||
for conv in self.convs:
|
||||
x = conv(x)
|
||||
return x.squeeze((-2, -1))
|
||||
|
||||
class MotionEncoder(torch.nn.Module):
|
||||
def __init__(self, dim=512, motion_dim=20):
|
||||
super().__init__()
|
||||
self.net_app = AppearanceEncoder(dim)
|
||||
self.fc = torch.nn.Sequential(*[EqualLinear(dim, dim) for _ in range(4)] + [EqualLinear(dim, motion_dim)])
|
||||
|
||||
def encode_motion(self, x):
|
||||
return self.fc(self.net_app(x))
|
||||
|
||||
class MotionProjector(torch.nn.Module):
|
||||
def __init__(self, m_dim):
|
||||
super().__init__()
|
||||
self.weight = torch.nn.Parameter(torch.randn(512, m_dim))
|
||||
self.motion_dim = m_dim
|
||||
|
||||
def forward(self, input):
|
||||
stabilized_weight = self.weight + 1e-8 * torch.eye(512, self.motion_dim, device=self.weight.device, dtype=self.weight.dtype)
|
||||
Q, _ = torch.linalg.qr(stabilized_weight)
|
||||
if input is None:
|
||||
return Q
|
||||
return torch.sum(input.unsqueeze(-1) * Q.T, dim=1)
|
||||
|
||||
class MotionDecoder(torch.nn.Module):
|
||||
def __init__(self, m_dim):
|
||||
super().__init__()
|
||||
self.direction = MotionProjector(m_dim)
|
||||
|
||||
class MotionExtractor(torch.nn.Module):
|
||||
def __init__(self, s_dim=512, m_dim=20):
|
||||
super().__init__()
|
||||
self.enc = MotionEncoder(s_dim, m_dim)
|
||||
self.dec = MotionDecoder(m_dim)
|
||||
|
||||
def forward(self, img):
|
||||
motion_feat = self.enc.encode_motion(img)
|
||||
return self.dec.direction(motion_feat)
|
||||
Reference in New Issue
Block a user