Merge branch 's2v'
This commit is contained in:
@@ -12,6 +12,7 @@ from .nodes_model_loading import NODE_CLASS_MAPPINGS as MODEL_LOADING_NODE_CLASS
|
||||
from .nodes_utility import NODE_CLASS_MAPPINGS as UTILITY_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UTILITY_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .cache_methods.nodes_cache import NODE_CLASS_MAPPINGS as NODE_CACHE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as NODE_CACHE_DISPLAY_NAME_MAPPINGS
|
||||
from .nodes_deprecated import NODE_CLASS_MAPPINGS as DEPRECATED_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .s2v.nodes import NODE_CLASS_MAPPINGS as S2V_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as S2V_NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
try:
|
||||
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS
|
||||
@@ -58,6 +59,7 @@ NODE_CLASS_MAPPINGS.update(NODE_CACHE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(QWEN_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(MTV_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(S2V_NODE_CLASS_MAPPINGS)
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
@@ -75,5 +77,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(NODE_CACHE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(S2V_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
@@ -261,6 +261,14 @@ class MultiTalkWav2VecEmbeds:
|
||||
offset += length
|
||||
multitalk_audio_features = full_list
|
||||
|
||||
# if audio_encoder_output is not None:
|
||||
# all_layers = audio_encoder_output["encoded_audio_all_layers"]
|
||||
# audio_feat = torch.stack(all_layers, dim=0).squeeze(1)[1:] # shape: [num_layers, T, 512]
|
||||
# audio_feat = audio_feat.movedim(0, 1)
|
||||
# print("audio_feat mean", audio_feat.mean())
|
||||
# print("audio_feat min max", audio_feat.min(), audio_feat.max())
|
||||
# multitalk_audio_features.append(audio_feat.cpu().detach())
|
||||
|
||||
# fallback
|
||||
if len(multitalk_audio_features) == 0:
|
||||
raise RuntimeError("No valid audio embeddings extracted, please check inputs")
|
||||
|
||||
@@ -1838,6 +1838,8 @@ class WanVideoSampler:
|
||||
log.info(f"sigmas: {sample_scheduler.sigmas}")
|
||||
else:
|
||||
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
|
||||
|
||||
log.info(f"timesteps: {timesteps}")
|
||||
total_steps = steps
|
||||
steps = len(timesteps)
|
||||
|
||||
@@ -2237,6 +2239,35 @@ class WanVideoSampler:
|
||||
log.info(f"mtv_motion_rotary_emb: {motion_rotary_emb[0].shape}")
|
||||
mtv_freqs = mtv_freqs.to(device, dtype)
|
||||
|
||||
#region S2V
|
||||
s2v_audio_input = s2v_ref_latent = s2v_pose = s2v_ref_motion = None
|
||||
framepack = False
|
||||
s2v_audio_embeds = image_embeds.get("audio_embeds", None)
|
||||
if s2v_audio_embeds is not None:
|
||||
log.info(f"Using S2V audio embeddings")
|
||||
framepack = s2v_audio_embeds.get("enable_framepack", False)
|
||||
if framepack and context_options is not None:
|
||||
raise ValueError("S2V framepack and context windows cannot be used at the same time")
|
||||
|
||||
s2v_audio_input = s2v_audio_embeds.get("audio_embed_bucket", None)
|
||||
if s2v_audio_input is not None:
|
||||
#s2v_audio_input = s2v_audio_input[..., 0:image_embeds["num_frames"]]
|
||||
s2v_audio_input = s2v_audio_input.to(device, dtype)
|
||||
s2v_audio_scale = s2v_audio_embeds["audio_scale"]
|
||||
s2v_ref_latent = s2v_audio_embeds.get("ref_latent", None)
|
||||
if s2v_ref_latent is not None:
|
||||
s2v_ref_latent = s2v_ref_latent.to(device, dtype)
|
||||
s2v_ref_motion = s2v_audio_embeds.get("ref_motion", None)
|
||||
if s2v_ref_motion is not None:
|
||||
s2v_ref_motion = s2v_ref_motion.to(device, dtype)
|
||||
s2v_pose = s2v_audio_embeds.get("pose_latent", None)
|
||||
if s2v_pose is not None:
|
||||
s2v_pose = s2v_pose.to(device, dtype)
|
||||
s2v_pose_start_percent = s2v_audio_embeds.get("pose_start_percent", 0.0)
|
||||
s2v_pose_end_percent = s2v_audio_embeds.get("pose_end_percent", 1.0)
|
||||
s2v_num_repeat = s2v_audio_embeds.get("num_repeat", 1)
|
||||
vae = s2v_audio_embeds.get("vae", None)
|
||||
|
||||
# vid2vid
|
||||
noise_mask=original_image=None
|
||||
if samples is not None and not multitalk_sampling:
|
||||
@@ -2516,7 +2547,7 @@ class WanVideoSampler:
|
||||
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
|
||||
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):
|
||||
mtv_motion_tokens=None, s2v_audio_input=None, s2v_ref_motion=None, s2v_motion_frames=[1, 0], s2v_pose=None):
|
||||
nonlocal transformer
|
||||
z = z.to(dtype)
|
||||
autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear)
|
||||
@@ -2661,7 +2692,11 @@ class WanVideoSampler:
|
||||
else:
|
||||
pcd_data_input = pcd_data
|
||||
|
||||
|
||||
if s2v_pose is not None:
|
||||
if not ((s2v_pose_start_percent <= current_step_percentage <= s2v_pose_end_percent) or \
|
||||
(s2v_pose_end_percent > 0 and idx == 0 and current_step_percentage >= s2v_pose_start_percent)):
|
||||
s2v_pose = None
|
||||
|
||||
base_params = {
|
||||
'seq_len': seq_len, # sequence length
|
||||
'device': device, # main device
|
||||
@@ -2694,6 +2729,12 @@ class WanVideoSampler:
|
||||
"mtv_motion_rotary_emb": mtv_motion_rotary_emb if mtv_input is not None else None, # MTV-Crafter RoPE
|
||||
"mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0, # MTV-Crafter scaling
|
||||
"mtv_freqs": mtv_freqs if mtv_input is not None else None, # MTV-Crafter extra RoPE freqs
|
||||
"s2v_audio_input": s2v_audio_input, # official speech-to-video audio input
|
||||
"s2v_ref_latent": s2v_ref_latent, # speech-to-video reference latent
|
||||
"s2v_ref_motion": s2v_ref_motion, # speech-to-video reference motion latent
|
||||
"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
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -2838,6 +2879,7 @@ class WanVideoSampler:
|
||||
noise_pred = noise_pred_uncond_scaled + cfg_scale * filtered_cond * alpha
|
||||
else:
|
||||
noise_pred = noise_pred_uncond_scaled + cfg_scale * (noise_pred_cond - noise_pred_uncond_scaled)
|
||||
del noise_pred_uncond_scaled, noise_pred_cond, noise_pred_uncond
|
||||
|
||||
|
||||
return noise_pred, [cache_state_cond, cache_state_uncond]
|
||||
@@ -2848,7 +2890,7 @@ class WanVideoSampler:
|
||||
from .latent_preview import prepare_callback #custom for tiny VAE previews
|
||||
callback = prepare_callback(patcher, len(timesteps))
|
||||
|
||||
if not multitalk_sampling:
|
||||
if not multitalk_sampling and not framepack:
|
||||
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")
|
||||
|
||||
@@ -3206,6 +3248,19 @@ class WanVideoSampler:
|
||||
log.info(f"context window: {c}")
|
||||
log.info(f"motion_token_indices: {start_token_index}-{end_token_index}")
|
||||
|
||||
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)
|
||||
center_indices = torch.clamp(center_indices, min=0, max=s2v_audio_input.shape[-1] - 1)
|
||||
partial_s2v_audio_input = s2v_audio_input[..., center_indices]
|
||||
|
||||
partial_s2v_pose = None
|
||||
if s2v_pose is not None:
|
||||
partial_s2v_pose = s2v_pose[:, :, c].to(device, dtype)
|
||||
|
||||
partial_add_cond = None
|
||||
if add_cond is not None:
|
||||
partial_add_cond = add_cond[:, :, c].to(device, dtype)
|
||||
@@ -3223,7 +3278,7 @@ class WanVideoSampler:
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
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)
|
||||
mtv_motion_tokens=partial_mtv_motion_tokens, s2v_audio_input=partial_s2v_audio_input, s2v_motion_frames=[1, 0], s2v_pose=partial_s2v_pose)
|
||||
|
||||
if cache_args is not None:
|
||||
self.window_tracker.cache_states[window_id] = new_teacache
|
||||
@@ -3231,7 +3286,7 @@ class WanVideoSampler:
|
||||
window_mask = create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=is_looped, window_type=context_options["fuse_method"])
|
||||
noise_pred[:, c] += noise_pred_context * window_mask
|
||||
counter[:, c] += window_mask
|
||||
context_pbar.update_absolute(step_start_progress + (i + 1) * fraction_per_context, steps)
|
||||
context_pbar.update_absolute(step_start_progress + (i + 1) * fraction_per_context, len(timesteps))
|
||||
noise_pred /= counter
|
||||
#region multitalk
|
||||
elif multitalk_sampling:
|
||||
@@ -3530,7 +3585,7 @@ class WanVideoSampler:
|
||||
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
latent_model_input, cfg[i], positive, text_embeds["negative_prompt_embeds"],
|
||||
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
|
||||
timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input)
|
||||
|
||||
@@ -3648,21 +3703,157 @@ class WanVideoSampler:
|
||||
except:
|
||||
pass
|
||||
return {"video": gen_video_samples.permute(1, 2, 3, 0)},
|
||||
|
||||
# region framepack loop
|
||||
elif framepack:
|
||||
framepack_out = []
|
||||
ref_motion_image = None
|
||||
#infer_frames = image_embeds["num_frames"]
|
||||
infer_frames = s2v_audio_embeds.get("frame_window_size", 80)
|
||||
motion_frames = infer_frames - 7 #73 default
|
||||
lat_motion_frames = (motion_frames + 3) // 4
|
||||
lat_target_frames = (infer_frames + 3 + motion_frames) // 4 - lat_motion_frames
|
||||
|
||||
step_iteration_count = 0
|
||||
total_frames = s2v_audio_input.shape[-1]
|
||||
|
||||
s2v_motion_frames = [motion_frames, lat_motion_frames]
|
||||
|
||||
noise = torch.randn( #C, T, H, W
|
||||
48 if is_5b else 16,
|
||||
lat_target_frames,
|
||||
target_shape[2],
|
||||
target_shape[3],
|
||||
dtype=torch.float32,
|
||||
generator=seed_g,
|
||||
device=torch.device("cpu"))
|
||||
|
||||
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
|
||||
|
||||
if ref_motion_image is None:
|
||||
ref_motion_image = torch.zeros(
|
||||
[1, 3, motion_frames, latent.shape[2]*vae_upscale_factor, latent.shape[3]*vae_upscale_factor],
|
||||
dtype=vae.dtype,
|
||||
device=device)
|
||||
videos_last_frames = ref_motion_image
|
||||
|
||||
if s2v_pose is not None:
|
||||
pose_cond_list = []
|
||||
for r in range(s2v_num_repeat):
|
||||
pose_start = r * (infer_frames // 4)
|
||||
pose_end = pose_start + (infer_frames // 4)
|
||||
|
||||
cond_lat = s2v_pose[:, :, pose_start:pose_end]
|
||||
|
||||
pad_len = (infer_frames // 4) - cond_lat.shape[2]
|
||||
if pad_len > 0:
|
||||
pad = -torch.ones(cond_lat.shape[0], cond_lat.shape[1], pad_len, cond_lat.shape[3], cond_lat.shape[4], device=cond_lat.device, dtype=cond_lat.dtype)
|
||||
cond_lat = torch.cat([cond_lat, pad], dim=2)
|
||||
pose_cond_list.append(cond_lat.cpu())
|
||||
|
||||
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
|
||||
right_idx = r * infer_frames + infer_frames
|
||||
|
||||
s2v_audio_input_slice = s2v_audio_input[..., left_idx:right_idx]
|
||||
if s2v_audio_input_slice.shape[-1] < (right_idx - left_idx):
|
||||
pad_len = (right_idx - left_idx) - s2v_audio_input_slice.shape[-1]
|
||||
pad_shape = list(s2v_audio_input_slice.shape)
|
||||
pad_shape[-1] = pad_len
|
||||
pad = torch.zeros(pad_shape, device=s2v_audio_input_slice.device, dtype=s2v_audio_input_slice.dtype)
|
||||
log.info(f"Padding s2v_audio_input_slice from {s2v_audio_input_slice.shape[-1]} to {right_idx - left_idx}")
|
||||
s2v_audio_input_slice = torch.cat([s2v_audio_input_slice, pad], dim=-1)
|
||||
|
||||
if ref_motion_image is not None:
|
||||
input_motion_latents = ref_motion.clone().unsqueeze(0)
|
||||
else:
|
||||
input_motion_latents = None
|
||||
|
||||
s2v_pose_slice = None
|
||||
if s2v_pose is not None:
|
||||
s2v_pose_slice = pose_cond_list[r].to(device)
|
||||
|
||||
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
|
||||
|
||||
latent = noise.to(device)
|
||||
for i, t in enumerate(tqdm(timesteps, desc=f"Sampling audio indices {left_idx}-{right_idx}", position=0)):
|
||||
latent_model_input = latent.to(device)
|
||||
timestep = torch.tensor([t]).to(device)
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
latent_model_input,
|
||||
cfg[idx],
|
||||
text_embeds["prompt_embeds"],
|
||||
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, mtv_motion_tokens=mtv_motion_tokens,
|
||||
s2v_audio_input=s2v_audio_input_slice, s2v_ref_motion=input_motion_latents, s2v_motion_frames=s2v_motion_frames, s2v_pose=s2v_pose_slice)
|
||||
|
||||
latent = sample_scheduler.step(
|
||||
noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0),
|
||||
**scheduler_step_args)[0].squeeze(0)
|
||||
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, s2v_num_repeat*(len(timesteps)))
|
||||
del callback_latent
|
||||
step_iteration_count += 1
|
||||
del latent_model_input, noise_pred
|
||||
|
||||
|
||||
vae.to(device)
|
||||
decode_latents = torch.cat([ref_motion.unsqueeze(0), latent.unsqueeze(0)], dim=2)
|
||||
image = vae.decode(decode_latents.to(device, vae.dtype), device=device, pbar=False)[0]
|
||||
del decode_latents
|
||||
image = image.unsqueeze(0)[:, :, -infer_frames:]
|
||||
if r == 0:
|
||||
image = image[:, :, 3:]
|
||||
|
||||
framepack_out.append(image.cpu())
|
||||
|
||||
overlap_frames_num = min(motion_frames, image.shape[2])
|
||||
|
||||
videos_last_frames = torch.cat([
|
||||
videos_last_frames[:, :, overlap_frames_num:],
|
||||
image[:, :, -overlap_frames_num:]], dim=2).to(device, vae.dtype)
|
||||
|
||||
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)
|
||||
|
||||
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},
|
||||
|
||||
#region normal inference
|
||||
else:
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
latent_model_input,
|
||||
cfg[idx],
|
||||
text_embeds["prompt_embeds"],
|
||||
cfg[idx], text_embeds["prompt_embeds"],
|
||||
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, mtv_motion_tokens=mtv_motion_tokens)
|
||||
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input)
|
||||
if bidirectional_sampling:
|
||||
noise_pred_flipped, self.cache_state = predict_with_cfg(
|
||||
latent_model_input_flipped,
|
||||
cfg[idx],
|
||||
text_embeds["prompt_embeds"],
|
||||
cfg[idx], text_embeds["prompt_embeds"],
|
||||
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, mtv_motion_tokens=mtv_motion_tokens,reverse_time=True)
|
||||
|
||||
+10
-3
@@ -731,7 +731,8 @@ 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", "audio_proj"}
|
||||
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding",
|
||||
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer"}
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
pbar = ProgressBar(param_count)
|
||||
cnt = 0
|
||||
@@ -1081,7 +1082,9 @@ class WanVideoModelLoader:
|
||||
ffn2_dim = sd["blocks.0.ffn.2.weight"].shape[1]
|
||||
|
||||
model_type = "t2v"
|
||||
if not "text_embedding.0.weight" in sd:
|
||||
if "audio_injector.injector.0.k.weight" in sd:
|
||||
model_type = "s2v"
|
||||
elif not "text_embedding.0.weight" in sd:
|
||||
model_type = "no_cross_attn" #minimaxremover
|
||||
elif "model_type.Wan2_1-FLF2V-14B-720P" in sd or "img_emb.emb_pos" in sd or "flf2v" in model.lower():
|
||||
model_type = "fl2v"
|
||||
@@ -1188,7 +1191,11 @@ class WanVideoModelLoader:
|
||||
"add_ref_conv": True if "ref_conv.weight" in sd else False,
|
||||
"in_dim_ref_conv": sd["ref_conv.weight"].shape[1] if "ref_conv.weight" in sd else None,
|
||||
"add_control_adapter": True if "control_adapter.conv.weight" in sd else False,
|
||||
"use_motion_attn": True if "blocks.0.motion_attn.k.weight" in sd else False
|
||||
"use_motion_attn": True if "blocks.0.motion_attn.k.weight" in sd else False,
|
||||
"enable_adain": True if "audio_injector.injector_adain_layers.0.linear.weight" in sd else False,
|
||||
"cond_dim": sd["cond_encoder.weight"].shape[1] if "cond_encoder.weight" in sd else 0,
|
||||
"zero_timestep": model_type == "s2v",
|
||||
|
||||
}
|
||||
|
||||
with init_empty_weights():
|
||||
|
||||
+44
-2
@@ -1,6 +1,7 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
from comfy.utils import common_upscale
|
||||
from .utils import log
|
||||
|
||||
try:
|
||||
from server import PromptServer
|
||||
@@ -408,6 +409,45 @@ class WanVideoSigmaToStep:
|
||||
def convert(self, sigma):
|
||||
return (sigma,)
|
||||
|
||||
class NormalizeAudioLoudness:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"audio": ("AUDIO",),
|
||||
"lufs": ("FLOAT", {"default": -23.0, "min": -100.0, "max": 0.0, "step": 0.1, "tool": "Loudness Units relative to Full Scale, higher LUFS values (closer to 0) mean louder audio. Lower LUFS values (more negative) mean quieter audio."}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO", )
|
||||
RETURN_NAMES = ("audio", )
|
||||
FUNCTION = "normalize"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def normalize(self, audio, lufs):
|
||||
audio_input = audio["waveform"]
|
||||
sample_rate = audio["sample_rate"]
|
||||
if audio_input.dim() == 3:
|
||||
audio_input = audio_input.squeeze(0)
|
||||
audio_input_np = audio_input.detach().transpose(0, 1).numpy().astype(np.float32)
|
||||
audio_input_np = np.ascontiguousarray(audio_input_np)
|
||||
normalized_audio = self.loudness_norm(audio_input_np, sr=sample_rate, lufs=lufs)
|
||||
|
||||
out_audio = {"waveform": torch.from_numpy(normalized_audio).transpose(0, 1).unsqueeze(0).float(), "sample_rate": sample_rate}
|
||||
|
||||
return (out_audio, )
|
||||
|
||||
def loudness_norm(self, audio_array, sr=16000, lufs=-23):
|
||||
try:
|
||||
import pyloudnorm
|
||||
except:
|
||||
raise ImportError("pyloudnorm package is not installed")
|
||||
meter = pyloudnorm.Meter(sr)
|
||||
loudness = meter.integrated_loudness(audio_array)
|
||||
if abs(loudness) > 100:
|
||||
return audio_array
|
||||
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
|
||||
return normalized_audio
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
|
||||
"WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame,
|
||||
@@ -416,7 +456,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"DummyComfyWanModelObject": DummyComfyWanModelObject,
|
||||
"WanVideoLatentReScale": WanVideoLatentReScale,
|
||||
"CreateScheduleFloatList": CreateScheduleFloatList,
|
||||
"WanVideoSigmaToStep": WanVideoSigmaToStep
|
||||
"WanVideoSigmaToStep": WanVideoSigmaToStep,
|
||||
"NormalizeAudioLoudness": NormalizeAudioLoudness
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
|
||||
@@ -426,5 +467,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DummyComfyWanModelObject": "Dummy Comfy Wan Model Object",
|
||||
"WanVideoLatentReScale": "WanVideo Latent ReScale",
|
||||
"CreateScheduleFloatList": "Create Schedule Float List",
|
||||
"WanVideoSigmaToStep": "WanVideo Sigma To Step"
|
||||
"WanVideoSigmaToStep": "WanVideo Sigma To Step",
|
||||
"NormalizeAudioLoudness": "Normalize Audio Loudness"
|
||||
}
|
||||
+184
@@ -0,0 +1,184 @@
|
||||
import folder_paths
|
||||
import math
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
|
||||
def get_sample_indices(original_fps,
|
||||
total_frames,
|
||||
target_fps,
|
||||
num_sample,
|
||||
fixed_start=None):
|
||||
required_duration = num_sample / target_fps
|
||||
required_origin_frames = int(np.ceil(required_duration * original_fps))
|
||||
if required_duration > total_frames / original_fps:
|
||||
raise ValueError("required_duration must be less than video length")
|
||||
|
||||
if not fixed_start is None and fixed_start >= 0:
|
||||
start_frame = fixed_start
|
||||
else:
|
||||
max_start = total_frames - required_origin_frames
|
||||
if max_start < 0:
|
||||
raise ValueError("video length is too short")
|
||||
start_frame = np.random.randint(0, max_start + 1)
|
||||
start_time = start_frame / original_fps
|
||||
|
||||
end_time = start_time + required_duration
|
||||
time_points = np.linspace(start_time, end_time, num_sample, endpoint=False)
|
||||
|
||||
frame_indices = np.round(np.array(time_points) * original_fps).astype(int)
|
||||
frame_indices = np.clip(frame_indices, 0, total_frames - 1)
|
||||
return frame_indices
|
||||
|
||||
def linear_interpolation(features, input_fps, output_fps, output_len=None):
|
||||
"""
|
||||
features: shape=[1, T, 512]
|
||||
input_fps: fps for audio, f_a
|
||||
output_fps: fps for video, f_m
|
||||
output_len: video length
|
||||
"""
|
||||
features = features.transpose(1, 2) # [1, 512, T]
|
||||
seq_len = features.shape[2] / float(input_fps) # T/f_a
|
||||
if output_len is None:
|
||||
output_len = int(seq_len * output_fps) # f_m*T/f_a
|
||||
output_features = F.interpolate(
|
||||
features, size=output_len, align_corners=True,
|
||||
mode='linear') # [1, 512, output_len]
|
||||
return output_features.transpose(1, 2) # [1, output_len, 512]
|
||||
|
||||
class WanVideoAddS2VEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"frame_window_size": ("INT", {"default": 80, "min": 1, "max": 100000, "step": 1, "tooltip": "Number of frames in a single window"}),
|
||||
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for audio embeddings"}),
|
||||
"pose_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage for pose embeddings"}),
|
||||
"pose_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage for pose embeddings"})
|
||||
},
|
||||
"optional": {
|
||||
"audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",),
|
||||
"ref_latent": ("LATENT",),
|
||||
"pose_latent": ("LATENT",),
|
||||
"vae": ("WANVAE",),
|
||||
"enable_framepack": ("BOOLEAN", {"default": False, "tooltip": "Enable Framepack sampling loop, not compatible with context windows"})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "INT",)
|
||||
RETURN_NAMES = ("image_embeds", "audio_frame_count")
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, frame_window_size, audio_encoder_output=None, audio_scale=1.0, ref_latent=None, pose_latent=None, vae=None, pose_start_percent=0.0, pose_end_percent=1.0, enable_framepack=False):
|
||||
audio_frame_count=0
|
||||
if audio_encoder_output is not None:
|
||||
all_layers = audio_encoder_output["encoded_audio_all_layers"]
|
||||
audio_feat = torch.stack(all_layers, dim=0).squeeze(1) # shape: [num_layers, T, 512]
|
||||
|
||||
print("audio_feat in", audio_feat.shape)
|
||||
input_fps = 50 # determined by the model itself
|
||||
output_fps = 30 # determined by the model itself
|
||||
bucket_fps = 16 # target fps for the generation
|
||||
|
||||
if input_fps != output_fps:
|
||||
audio_feat = linear_interpolation(audio_feat, input_fps=input_fps, output_fps=output_fps)
|
||||
print("audio_feat after interpolation", audio_feat.shape)
|
||||
|
||||
audio_feat = audio_feat[:, :embeds["num_frames"] * output_fps // bucket_fps, :]
|
||||
print("audio_feat after trim", audio_feat.shape)
|
||||
|
||||
self.video_rate = output_fps
|
||||
|
||||
audio_embed_bucket, num_repeat = self.get_audio_embed_bucket_fps(
|
||||
audio_feat,
|
||||
fps=bucket_fps,
|
||||
batch_frames=frame_window_size
|
||||
)
|
||||
print("audio_embed_bucket", audio_embed_bucket.shape)
|
||||
|
||||
audio_embed_bucket = audio_embed_bucket.unsqueeze(0)
|
||||
if len(audio_embed_bucket.shape) == 3:
|
||||
audio_embed_bucket = audio_embed_bucket.permute(0, 2, 1)
|
||||
elif len(audio_embed_bucket.shape) == 4:
|
||||
audio_embed_bucket = audio_embed_bucket.permute(0, 2, 3, 1)
|
||||
|
||||
audio_frame_count = audio_embed_bucket.shape[-1]
|
||||
|
||||
print("audio_embed_bucket", audio_embed_bucket.shape)
|
||||
|
||||
new_entry = {
|
||||
"audio_embed_bucket": audio_embed_bucket if audio_encoder_output is not None else None,
|
||||
"num_repeat": num_repeat if audio_encoder_output is not None else None,
|
||||
"ref_latent": ref_latent["samples"] if ref_latent is not None else None,
|
||||
"pose_latent": pose_latent["samples"] if pose_latent is not None else None,
|
||||
"audio_scale": audio_scale,
|
||||
"vae": vae,
|
||||
"pose_start_percent": pose_start_percent,
|
||||
"pose_end_percent": pose_end_percent,
|
||||
"enable_framepack": enable_framepack,
|
||||
"frame_window_size": frame_window_size
|
||||
}
|
||||
updated = dict(embeds)
|
||||
updated["audio_embeds"] = new_entry
|
||||
return (updated, audio_frame_count)
|
||||
|
||||
def get_audio_embed_bucket_fps(self, audio_embed, fps=16, batch_frames=81, m=0):
|
||||
num_layers, audio_frame_num, audio_dim = audio_embed.shape
|
||||
|
||||
if num_layers > 1:
|
||||
return_all_layers = True
|
||||
else:
|
||||
return_all_layers = False
|
||||
|
||||
scale = self.video_rate / fps
|
||||
|
||||
min_batch_num = int(audio_frame_num / (batch_frames * scale)) + 1
|
||||
|
||||
bucket_num = min_batch_num * batch_frames
|
||||
padd_audio_num = math.ceil(min_batch_num * batch_frames / fps * self.video_rate) - audio_frame_num
|
||||
batch_idx = get_sample_indices(
|
||||
original_fps=self.video_rate,
|
||||
total_frames=audio_frame_num + padd_audio_num,
|
||||
target_fps=fps,
|
||||
num_sample=bucket_num,
|
||||
fixed_start=0)
|
||||
batch_audio_eb = []
|
||||
audio_sample_stride = int(self.video_rate / fps)
|
||||
for bi in batch_idx:
|
||||
if bi < audio_frame_num:
|
||||
|
||||
chosen_idx = list(
|
||||
range(bi - m * audio_sample_stride,
|
||||
bi + (m + 1) * audio_sample_stride,
|
||||
audio_sample_stride))
|
||||
chosen_idx = [0 if c < 0 else c for c in chosen_idx]
|
||||
chosen_idx = [
|
||||
audio_frame_num - 1 if c >= audio_frame_num else c
|
||||
for c in chosen_idx
|
||||
]
|
||||
|
||||
if return_all_layers:
|
||||
frame_audio_embed = audio_embed[:, chosen_idx].flatten(
|
||||
start_dim=-2, end_dim=-1)
|
||||
else:
|
||||
frame_audio_embed = audio_embed[0][chosen_idx].flatten()
|
||||
else:
|
||||
frame_audio_embed = \
|
||||
torch.zeros([audio_dim * (2 * m + 1)], device=audio_embed.device) if not return_all_layers \
|
||||
else torch.zeros([num_layers, audio_dim * (2 * m + 1)], device=audio_embed.device)
|
||||
batch_audio_eb.append(frame_audio_embed)
|
||||
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb],
|
||||
dim=0)
|
||||
|
||||
return batch_audio_eb, min_batch_num
|
||||
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddS2VEmbeds": WanVideoAddS2VEmbeds,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddS2VEmbeds": "WanVideo Add S2V Embeds",
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+66
-40
@@ -19,6 +19,11 @@ import comfy.model_management as mm
|
||||
from comfy.utils import load_torch_file, ProgressBar, common_upscale
|
||||
from comfy.clip_vision import clip_preprocess, ClipVisionModel
|
||||
from comfy.cli_args import args, LatentPreviewMethod
|
||||
from ..nodes_model_loading import load_weights
|
||||
from ..nodes import offload_transformer
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
@@ -142,33 +147,51 @@ class WanVideoDiffusionForcingSampler:
|
||||
patcher = model
|
||||
model = model.model
|
||||
transformer = model.diffusion_model
|
||||
dtype = model["dtype"]
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
dtype = model["base_dtype"]
|
||||
weight_dtype = model["weight_dtype"]
|
||||
fp8_matmul = model["fp8_matmul"]
|
||||
gguf = model["gguf"]
|
||||
gguf_reader = model["gguf_reader"]
|
||||
control_lora = model["control_lora"]
|
||||
|
||||
transformer_options = patcher.model_options.get("transformer_options", None)
|
||||
merge_loras = transformer_options["merge_loras"]
|
||||
|
||||
patch_linear = transformer_options.get("patch_linear", False)
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
|
||||
transformer.blocks_to_swap = block_swap_args.get("blocks_to_swap", 0)
|
||||
transformer.vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", 0)
|
||||
transformer.prefetch_blocks = block_swap_args.get("prefetch_blocks", 0)
|
||||
transformer.block_swap_debug = block_swap_args.get("block_swap_debug", False)
|
||||
transformer.offload_img_emb = block_swap_args.get("offload_img_emb", False)
|
||||
transformer.offload_txt_emb = block_swap_args.get("offload_txt_emb", False)
|
||||
|
||||
if gguf:
|
||||
is_5b = transformer.out_dim == 48
|
||||
vae_upscale_factor = 16 if is_5b else 8
|
||||
|
||||
# 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)
|
||||
|
||||
if 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)
|
||||
set_lora_params_gguf(transformer, patcher.patches)
|
||||
elif len(patcher.patches) != 0 and patch_linear:
|
||||
transformer.patched_linear = True
|
||||
elif len(patcher.patches) != 0 and transformer.patched_linear: #handle patched linear layers (unmerged loras, fp8 scaled)
|
||||
log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model")
|
||||
if not merge_loras and fp8_matmul:
|
||||
raise NotImplementedError("FP8 matmul with unmerged LoRAs is not supported")
|
||||
set_lora_params(transformer, patcher.patches)
|
||||
else:
|
||||
remove_lora_from_module(transformer)
|
||||
remove_lora_from_module(transformer) #clear possible unmerged lora weights
|
||||
|
||||
transformer.lora_scheduling_enabled = transformer_options.get("lora_scheduling_enabled", False)
|
||||
|
||||
#torch.compile
|
||||
if model["auto_cpu_offload"] is False:
|
||||
transformer = compile_model(transformer, model["compile_args"])
|
||||
|
||||
|
||||
steps = int(steps/denoise_strength)
|
||||
|
||||
timesteps = None
|
||||
@@ -367,34 +390,39 @@ class WanVideoDiffusionForcingSampler:
|
||||
callback = prepare_callback(patcher, steps)
|
||||
|
||||
#blockswap init
|
||||
if transformer_options is not None:
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
#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)
|
||||
|
||||
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)
|
||||
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()
|
||||
elif model["manual_offloading"]:
|
||||
transformer.to(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)
|
||||
|
||||
# Initialize Cache if enabled
|
||||
transformer.enable_teacache = transformer.enable_magcache = False
|
||||
@@ -610,10 +638,8 @@ class WanVideoDiffusionForcingSampler:
|
||||
transformer.teacache_state.clear_all()
|
||||
|
||||
if force_offload:
|
||||
if model["manual_offloading"]:
|
||||
transformer.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
if not model["auto_cpu_offload"]:
|
||||
offload_transformer(transformer)
|
||||
|
||||
try:
|
||||
print_memory(device)
|
||||
|
||||
+503
-110
@@ -20,7 +20,7 @@ except:
|
||||
|
||||
from .attention import attention
|
||||
import numpy as np
|
||||
|
||||
from copy import deepcopy
|
||||
from tqdm import tqdm
|
||||
import gc
|
||||
|
||||
@@ -31,11 +31,94 @@ from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
|
||||
|
||||
from ...MTV.mtv import apply_rotary_emb
|
||||
|
||||
#from .s2v.motioner import MotionerTransformers, FramePackMotioner, rope_precompute
|
||||
|
||||
#from comfy.ldm.wan.model import FramePackMotioner
|
||||
class FramePackMotioner(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
inner_dim=1024,
|
||||
num_heads=16, # Used to indicate the number of heads in the backbone network; unrelated to this module's design
|
||||
zip_frame_buckets=[1, 2, 16], # Three numbers representing the number of frames sampled for patch operations from the nearest to the farthest frames
|
||||
drop_mode="drop", # If not "drop", it will use "padd", meaning padding instead of deletion
|
||||
):
|
||||
super().__init__()
|
||||
self.proj = nn.Conv3d(16, inner_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2))
|
||||
self.proj_2x = nn.Conv3d(16, inner_dim, kernel_size=(2, 4, 4), stride=(2, 4, 4))
|
||||
self.proj_4x = nn.Conv3d(16, inner_dim, kernel_size=(4, 8, 8), stride=(4, 8, 8))
|
||||
self.zip_frame_buckets = zip_frame_buckets
|
||||
|
||||
self.inner_dim = inner_dim
|
||||
self.num_heads = num_heads
|
||||
self.drop_mode = drop_mode
|
||||
|
||||
def forward(self, motion_latents, rope_embedder, add_last_motion=2):
|
||||
lat_height, lat_width = motion_latents.shape[3], motion_latents.shape[4]
|
||||
padd_lat = torch.zeros(motion_latents.shape[0], 16, sum(self.zip_frame_buckets), lat_height, lat_width).to(device=motion_latents.device, dtype=motion_latents.dtype)
|
||||
overlap_frame = min(padd_lat.shape[2], motion_latents.shape[2])
|
||||
if overlap_frame > 0:
|
||||
padd_lat[:, :, -overlap_frame:] = motion_latents[:, :, -overlap_frame:]
|
||||
|
||||
if add_last_motion < 2 and self.drop_mode != "drop":
|
||||
zero_end_frame = sum(self.zip_frame_buckets[:len(self.zip_frame_buckets) - add_last_motion - 1])
|
||||
padd_lat[:, :, -zero_end_frame:] = 0
|
||||
|
||||
clean_latents_4x, clean_latents_2x, clean_latents_post = padd_lat[:, :, -sum(self.zip_frame_buckets):, :, :].split(self.zip_frame_buckets[::-1], dim=2) # 16, 2 ,1
|
||||
|
||||
# patchfy
|
||||
clean_latents_post = self.proj(clean_latents_post).flatten(2).transpose(1, 2)
|
||||
clean_latents_2x = self.proj_2x(clean_latents_2x)
|
||||
l_2x_shape = clean_latents_2x.shape
|
||||
clean_latents_2x = clean_latents_2x.flatten(2).transpose(1, 2)
|
||||
clean_latents_4x = self.proj_4x(clean_latents_4x)
|
||||
l_4x_shape = clean_latents_4x.shape
|
||||
clean_latents_4x = clean_latents_4x.flatten(2).transpose(1, 2)
|
||||
|
||||
if add_last_motion < 2 and self.drop_mode == "drop":
|
||||
clean_latents_post = clean_latents_post[:, :0] if add_last_motion < 2 else clean_latents_post
|
||||
clean_latents_2x = clean_latents_2x[:, :0] if add_last_motion < 1 else clean_latents_2x
|
||||
|
||||
motion_lat = torch.cat([clean_latents_post, clean_latents_2x, clean_latents_4x], dim=1)
|
||||
|
||||
rope_post = rope_embedder.rope_encode_comfy(1, lat_height, lat_width, t_start=-1, device=motion_latents.device, dtype=motion_latents.dtype)
|
||||
rope_2x = rope_embedder.rope_encode_comfy(1, lat_height, lat_width, t_start=-3, steps_h=l_2x_shape[-2], steps_w=l_2x_shape[-1], device=motion_latents.device, dtype=motion_latents.dtype)
|
||||
rope_4x = rope_embedder.rope_encode_comfy(4, lat_height, lat_width, t_start=-19, steps_h=l_4x_shape[-2], steps_w=l_4x_shape[-1], device=motion_latents.device, dtype=motion_latents.dtype)
|
||||
|
||||
rope = torch.cat([rope_post, rope_2x, rope_4x], dim=1)
|
||||
return motion_lat, rope
|
||||
|
||||
from diffusers.models.attention import AdaLayerNorm
|
||||
|
||||
__all__ = ['WanModel']
|
||||
|
||||
from comfy import model_management as mm
|
||||
|
||||
|
||||
def zero_module(module):
|
||||
"""
|
||||
Zero out the parameters of a module and return it.
|
||||
"""
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
|
||||
|
||||
def torch_dfs(model: nn.Module, parent_name='root'):
|
||||
module_names, modules = [], []
|
||||
current_name = parent_name if parent_name else 'root'
|
||||
module_names.append(current_name)
|
||||
modules.append(model)
|
||||
|
||||
for name, child in model.named_children():
|
||||
if parent_name:
|
||||
child_name = f'{parent_name}.{name}'
|
||||
else:
|
||||
child_name = name
|
||||
child_modules, child_names = torch_dfs(child, child_name)
|
||||
module_names += child_names
|
||||
modules += child_modules
|
||||
return modules, module_names
|
||||
|
||||
#from comfy.ldm.flux.math import apply_rope as apply_rope_comfy
|
||||
def apply_rope_comfy(xq, xk, freqs_cis):
|
||||
xq_ = xq.to(dtype=freqs_cis.dtype).reshape(*xq.shape[:-1], -1, 1, 2)
|
||||
@@ -165,7 +248,6 @@ def sinusoidal_embedding_1d(dim, position):
|
||||
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
||||
return x
|
||||
|
||||
|
||||
def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0):
|
||||
assert dim % 2 == 0
|
||||
exponents = torch.arange(0, dim, 2, dtype=torch.float64).div(dim)
|
||||
@@ -545,7 +627,7 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
|
||||
num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
|
||||
inner_t=None, inner_c=None, cross_freqs=None,
|
||||
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, **kwargs):
|
||||
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
# compute query
|
||||
q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d)
|
||||
@@ -583,19 +665,20 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
# FantasyPortrait adapter attention
|
||||
if adapter_proj is not None:
|
||||
if len(adapter_proj.shape) == 4:
|
||||
adapter_q = q.view(b * num_latent_frames, -1, n, d)
|
||||
q_in = q[:, :orig_seq_len]
|
||||
adapter_q = q_in.view(b * num_latent_frames, -1, n, d)
|
||||
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
|
||||
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
|
||||
|
||||
adapter_x = attention(adapter_q, ip_key, ip_value, attention_mode=self.attention_mode)
|
||||
adapter_x = adapter_x.view(b, q.size(1), n, d)
|
||||
adapter_x = adapter_x.view(b, q_in.size(1), n, d)
|
||||
adapter_x = adapter_x.flatten(2)
|
||||
elif len(adapter_proj.shape) == 3:
|
||||
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b, -1, n, d)
|
||||
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b, -1, n, d)
|
||||
adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode)
|
||||
adapter_x = attention(q_in, ip_key, ip_value, attention_mode=self.attention_mode)
|
||||
adapter_x = adapter_x.flatten(2)
|
||||
x = x + adapter_x * ip_scale
|
||||
x[:, :orig_seq_len] = x[:, :orig_seq_len] + adapter_x * ip_scale
|
||||
|
||||
return self.o(x)
|
||||
|
||||
@@ -611,7 +694,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None,
|
||||
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
|
||||
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, **kwargs):
|
||||
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
@@ -761,6 +844,7 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
# modulation
|
||||
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
|
||||
self.seg_idx = None
|
||||
|
||||
@torch.compiler.disable()
|
||||
def get_mod(self, e):
|
||||
@@ -770,8 +854,24 @@ class WanAttentionBlock(nn.Module):
|
||||
e = (self.modulation.unsqueeze(2) + e).chunk(6, dim=1) # 1, 6, 1, dim
|
||||
return [ei.squeeze(1) for ei in e]
|
||||
|
||||
def modulate(self, x, shift_msa, scale_msa):
|
||||
return torch.addcmul(shift_msa, x, 1 + scale_msa)
|
||||
def modulate(self, x, shift_msa, scale_msa, seg_idx=None):
|
||||
"""
|
||||
Modulate x with shift and scale. If seg_idx is provided, apply segmented modulation.
|
||||
"""
|
||||
norm_x = self.norm1(x)
|
||||
if seg_idx is not None:
|
||||
parts = []
|
||||
for i in range(2):
|
||||
part = torch.addcmul(
|
||||
shift_msa[:, i:i + 1],
|
||||
norm_x[:, seg_idx[i]:seg_idx[i + 1]],
|
||||
1 + scale_msa[:, i:i + 1]
|
||||
)
|
||||
parts.append(part)
|
||||
norm_x = torch.cat(parts, dim=1)
|
||||
return norm_x
|
||||
else:
|
||||
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
|
||||
|
||||
def ffn_chunked(self, x, shift_mlp, scale_mlp, num_chunks=4):
|
||||
modulated_input = torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp)
|
||||
@@ -808,6 +908,7 @@ class WanAttentionBlock(nn.Module):
|
||||
audio_proj=None,
|
||||
audio_scale=1.0,
|
||||
num_latent_frames=21,
|
||||
original_seq_len=None,
|
||||
enhance_enabled=False,
|
||||
block_mask=None,
|
||||
nag_params={},
|
||||
@@ -838,9 +939,16 @@ class WanAttentionBlock(nn.Module):
|
||||
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
|
||||
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
|
||||
"""
|
||||
#e = (self.modulation.to(e.device) + e).chunk(6, dim=1)
|
||||
self.original_seq_len = original_seq_len
|
||||
self.zero_timestep = len(e) == 2
|
||||
if self.zero_timestep: #s2v zero timestep
|
||||
self.seg_idx = e[1]
|
||||
self.seg_idx = min(max(0, self.seg_idx), x.size(1))
|
||||
self.seg_idx = [0, self.seg_idx, x.size(1)]
|
||||
e = e[0]
|
||||
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device))
|
||||
input_x = self.modulate(self.norm1(x), shift_msa, scale_msa)
|
||||
input_x = self.modulate(x, shift_msa, scale_msa, seg_idx=self.seg_idx)
|
||||
|
||||
if x_ip is not None:
|
||||
shift_msa_ip, scale_msa_ip, gate_msa_ip, shift_mlp_ip, scale_mlp_ip, gate_mlp_ip = self.get_mod(e_ip.to(x.device))
|
||||
@@ -958,7 +1066,14 @@ class WanAttentionBlock(nn.Module):
|
||||
y[:, -self.cond_size :],
|
||||
)
|
||||
|
||||
x = x.addcmul(y, gate_msa)
|
||||
if self.zero_timestep:
|
||||
z = []
|
||||
for i in range(2):
|
||||
z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_msa[:, i:i + 1])
|
||||
y = torch.cat(z, dim=1)
|
||||
x = x.add(y)
|
||||
else:
|
||||
x = x.addcmul(y, gate_msa)
|
||||
|
||||
# cross-attention & ffn function
|
||||
if context is not None:
|
||||
@@ -996,7 +1111,7 @@ class WanAttentionBlock(nn.Module):
|
||||
audio_proj=audio_proj, audio_scale=audio_scale,
|
||||
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
|
||||
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale)
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=self.original_seq_len)
|
||||
# MultiTalk
|
||||
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
|
||||
x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding,
|
||||
@@ -1008,11 +1123,27 @@ class WanAttentionBlock(nn.Module):
|
||||
x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
|
||||
x = x + x_motion * mtv_strength
|
||||
|
||||
if self.rope_func == "comfy_chunked":
|
||||
if self.rope_func == "comfy_chunked" and not self.zero_timestep:
|
||||
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
||||
else:
|
||||
y = self.ffn(torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp))
|
||||
x = x.addcmul(y, gate_mlp)
|
||||
norm2_x = self.norm2(x)
|
||||
if self.zero_timestep:
|
||||
parts = []
|
||||
for i in range(2):
|
||||
parts.append(norm2_x[:, self.seg_idx[i]:self.seg_idx[i + 1]] *
|
||||
(1 + scale_mlp[:, i:i + 1]) + shift_mlp[:, i:i + 1])
|
||||
norm2_x = torch.cat(parts, dim=1)
|
||||
y = self.ffn(norm2_x)
|
||||
else:
|
||||
y = self.ffn(torch.addcmul(shift_mlp, norm2_x, 1 + scale_mlp))
|
||||
if self.zero_timestep:
|
||||
z = []
|
||||
for i in range(2):
|
||||
z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_mlp[:, i:i + 1])
|
||||
y = torch.cat(z, dim=1)
|
||||
x = x.add(y)
|
||||
else:
|
||||
x = x.addcmul(y, gate_mlp)
|
||||
return x
|
||||
|
||||
@torch.compiler.disable()
|
||||
@@ -1176,41 +1307,145 @@ class MLPProj(torch.nn.Module):
|
||||
clip_extra_context_tokens = self.proj(image_embeds)
|
||||
return clip_extra_context_tokens
|
||||
|
||||
from .s2v.auxi_blocks import MotionEncoder_tc
|
||||
|
||||
|
||||
class CausalAudioEncoder(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim=5120,
|
||||
num_layers=25,
|
||||
out_dim=2048,
|
||||
video_rate=8,
|
||||
num_token=4,
|
||||
need_global=False):
|
||||
super().__init__()
|
||||
self.encoder = MotionEncoder_tc(
|
||||
in_dim=dim,
|
||||
hidden_dim=out_dim,
|
||||
num_heads=num_token,
|
||||
need_global=need_global)
|
||||
weight = torch.ones((1, num_layers, 1, 1)) * 0.01
|
||||
|
||||
self.weights = torch.nn.Parameter(weight)
|
||||
self.act = torch.nn.SiLU()
|
||||
|
||||
def forward(self, features):
|
||||
# features B * num_layers * dim * video_length
|
||||
weights = self.act(self.weights)
|
||||
weights_sum = weights.sum(dim=1, keepdims=True)
|
||||
weighted_feat = ((features * weights) / weights_sum).sum(
|
||||
dim=1) # b dim f
|
||||
weighted_feat = weighted_feat.permute(0, 2, 1) # b f dim
|
||||
res = self.encoder(weighted_feat) # b f n dim
|
||||
|
||||
return res # b f n dim
|
||||
|
||||
|
||||
class AudioCrossAttention(WanT2VCrossAttention):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
|
||||
class AudioInjector_WAN(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
all_modules,
|
||||
all_modules_names,
|
||||
dim=2048,
|
||||
num_heads=32,
|
||||
inject_layer=[0, 27],
|
||||
root_net=None,
|
||||
enable_adain=False,
|
||||
adain_dim=2048,
|
||||
need_adain_ont=False,
|
||||
attention_mode='sdpa'):
|
||||
super().__init__()
|
||||
self.injected_block_id = {}
|
||||
audio_injector_id = 0
|
||||
for mod_name, mod in zip(all_modules_names, all_modules):
|
||||
if isinstance(mod, WanAttentionBlock):
|
||||
for inject_id in inject_layer:
|
||||
if f'transformer_blocks.{inject_id}' in mod_name:
|
||||
self.injected_block_id[inject_id] = audio_injector_id
|
||||
audio_injector_id += 1
|
||||
|
||||
self.injector = nn.ModuleList([
|
||||
AudioCrossAttention(
|
||||
in_features=dim,
|
||||
out_features=dim,
|
||||
num_heads=num_heads,
|
||||
qk_norm=True,
|
||||
attention_mode=attention_mode
|
||||
) for _ in range(audio_injector_id)
|
||||
])
|
||||
self.injector_pre_norm_feat = nn.ModuleList([
|
||||
nn.LayerNorm(
|
||||
dim,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
) for _ in range(audio_injector_id)
|
||||
])
|
||||
self.injector_pre_norm_vec = nn.ModuleList([
|
||||
nn.LayerNorm(
|
||||
dim,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
) for _ in range(audio_injector_id)
|
||||
])
|
||||
if enable_adain:
|
||||
self.injector_adain_layers = nn.ModuleList([
|
||||
AdaLayerNorm(
|
||||
output_dim=dim * 2, embedding_dim=adain_dim, chunk_dim=1)
|
||||
for _ in range(audio_injector_id)
|
||||
])
|
||||
if need_adain_ont:
|
||||
self.injector_adain_output_layers = nn.ModuleList(
|
||||
[nn.Linear(dim, dim) for _ in range(audio_injector_id)])
|
||||
|
||||
class WanModel(torch.nn.Module):
|
||||
def __init__(self,
|
||||
model_type='t2v',
|
||||
patch_size=(1, 2, 2),
|
||||
text_len=512,
|
||||
in_dim=16,
|
||||
dim=2048,
|
||||
in_features=5120,
|
||||
out_features=5120,
|
||||
ffn_dim=8192,
|
||||
ffn2_dim=8192,
|
||||
freq_dim=256,
|
||||
text_dim=4096,
|
||||
out_dim=16,
|
||||
num_heads=16,
|
||||
num_layers=32,
|
||||
qk_norm=True,
|
||||
cross_attn_norm=True,
|
||||
eps=1e-6,
|
||||
attention_mode='sdpa',
|
||||
rope_func='comfy',
|
||||
main_device=torch.device('cuda'),
|
||||
offload_device=torch.device('cpu'),
|
||||
teacache_coefficients=[],
|
||||
magcache_ratios=[],
|
||||
vace_layers=None,
|
||||
vace_in_dim=None,
|
||||
inject_sample_info=False,
|
||||
add_ref_conv=False,
|
||||
in_dim_ref_conv=16,
|
||||
add_control_adapter=False,
|
||||
in_dim_control_adapter=24,
|
||||
use_motion_attn=False
|
||||
):
|
||||
model_type='t2v',
|
||||
patch_size=(1, 2, 2),
|
||||
text_len=512,
|
||||
in_dim=16,
|
||||
dim=2048,
|
||||
in_features=5120,
|
||||
out_features=5120,
|
||||
ffn_dim=8192,
|
||||
ffn2_dim=8192,
|
||||
freq_dim=256,
|
||||
text_dim=4096,
|
||||
out_dim=16,
|
||||
num_heads=16,
|
||||
num_layers=32,
|
||||
qk_norm=True,
|
||||
cross_attn_norm=True,
|
||||
eps=1e-6,
|
||||
attention_mode='sdpa',
|
||||
rope_func='comfy',
|
||||
main_device=torch.device('cuda'),
|
||||
offload_device=torch.device('cpu'),
|
||||
teacache_coefficients=[],
|
||||
magcache_ratios=[],
|
||||
vace_layers=None,
|
||||
vace_in_dim=None,
|
||||
inject_sample_info=False,
|
||||
add_ref_conv=False,
|
||||
in_dim_ref_conv=16,
|
||||
add_control_adapter=False,
|
||||
in_dim_control_adapter=24,
|
||||
use_motion_attn=False,
|
||||
#s2v
|
||||
cond_dim=0,
|
||||
audio_dim=1024,
|
||||
num_audio_token=4,
|
||||
enable_adain=False,
|
||||
adain_mode="attn_norm",
|
||||
audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39],
|
||||
zero_timestep=False
|
||||
):
|
||||
r"""
|
||||
Initialize the diffusion model backbone.
|
||||
|
||||
@@ -1361,7 +1596,7 @@ class WanModel(torch.nn.Module):
|
||||
])
|
||||
else:
|
||||
# blocks
|
||||
if model_type == 't2v':
|
||||
if model_type == 't2v' or model_type == 's2v':
|
||||
cross_attn_type = 't2v_cross_attn'
|
||||
elif model_type == 'i2v' or model_type == 'fl2v':
|
||||
cross_attn_type = 'i2v_cross_attn'
|
||||
@@ -1416,6 +1651,47 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
self.block_mask=None
|
||||
|
||||
#S2V
|
||||
self.zero_timestep = self.audio_injector = self.trainable_cond_mask =None
|
||||
if cond_dim > 0:
|
||||
self.cond_encoder = nn.Conv3d(
|
||||
cond_dim,
|
||||
self.dim,
|
||||
kernel_size=self.patch_size,
|
||||
stride=self.patch_size)
|
||||
if self.model_type == 's2v':
|
||||
self.enable_adain = enable_adain
|
||||
self.casual_audio_encoder = CausalAudioEncoder(
|
||||
dim=audio_dim,
|
||||
out_dim=self.dim,
|
||||
num_token=num_audio_token,
|
||||
need_global=enable_adain)
|
||||
all_modules, all_modules_names = torch_dfs(
|
||||
self.blocks, parent_name="root.transformer_blocks")
|
||||
self.audio_injector = AudioInjector_WAN(
|
||||
all_modules,
|
||||
all_modules_names,
|
||||
dim=self.dim,
|
||||
num_heads=self.num_heads,
|
||||
inject_layer=audio_inject_layers,
|
||||
root_net=self,
|
||||
enable_adain=enable_adain,
|
||||
adain_dim=self.dim,
|
||||
need_adain_ont=adain_mode != "attn_norm",
|
||||
attention_mode=attention_mode
|
||||
)
|
||||
self.trainable_cond_mask = nn.Embedding(3, self.dim)
|
||||
|
||||
self.frame_packer = FramePackMotioner(
|
||||
inner_dim=self.dim,
|
||||
num_heads=self.num_heads,
|
||||
zip_frame_buckets=[1, 2, 16],
|
||||
drop_mode='padd')
|
||||
self.adain_mode = adain_mode
|
||||
self.zero_timestep = zero_timestep
|
||||
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _prepare_blockwise_causal_attn_mask(
|
||||
device: torch.device | str, num_frames: int = 21,
|
||||
@@ -1566,6 +1842,77 @@ class WanModel(torch.nn.Module):
|
||||
block.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
return hints
|
||||
|
||||
def audio_injector_forward(self, block_idx, x, audio_emb, scale=1.0):
|
||||
if block_idx in self.audio_injector.injected_block_id.keys():
|
||||
audio_attn_id = self.audio_injector.injected_block_id[block_idx]
|
||||
num_frames = audio_emb.shape[1]# b f n c
|
||||
|
||||
input_x = x[:, :self.original_seq_len].clone() # b (f h w) c
|
||||
input_x = rearrange(input_x, "b (t n) c -> (b t) n c", t=num_frames)
|
||||
|
||||
if self.enable_adain and self.adain_mode == "attn_norm":
|
||||
audio_emb_global = self.audio_emb_global
|
||||
audio_emb_global = rearrange(audio_emb_global,"b t n c -> (b t) n c")
|
||||
attn_x = self.audio_injector.injector_adain_layers[audio_attn_id](input_x, temb=audio_emb_global[:, 0])
|
||||
else:
|
||||
attn_x = self.audio_injector.injector_pre_norm_feat[audio_attn_id](input_x)
|
||||
|
||||
attn_audio_emb = rearrange(audio_emb, "b t n c -> (b t) n c", t=num_frames)
|
||||
residual_out = self.audio_injector.injector[audio_attn_id](
|
||||
x=attn_x ,
|
||||
context=attn_audio_emb * scale,
|
||||
)
|
||||
residual_out = rearrange(residual_out, "(b t) n c -> b (t n) c", t=num_frames)
|
||||
x[:, :self.original_seq_len].add_(residual_out)
|
||||
|
||||
return x
|
||||
|
||||
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond=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])
|
||||
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
|
||||
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
|
||||
|
||||
if steps_t is None:
|
||||
steps_t = t_len
|
||||
if steps_h is None:
|
||||
steps_h = h_len
|
||||
if steps_w is None:
|
||||
steps_w = w_len
|
||||
|
||||
img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype)
|
||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
|
||||
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1)
|
||||
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1)
|
||||
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
|
||||
if attn_cond is not None:
|
||||
F_cond, H_cond, W_cond = attn_cond.shape[2], attn_cond.shape[3], attn_cond.shape[4]
|
||||
cond_f_len = ((F_cond + (self.patch_size[0] // 2)) // self.patch_size[0])
|
||||
cond_h_len = ((H_cond + (self.patch_size[1] // 2)) // self.patch_size[1])
|
||||
cond_w_len = ((W_cond + (self.patch_size[2] // 2)) // self.patch_size[2])
|
||||
cond_img_ids = torch.zeros((cond_f_len, cond_h_len, cond_w_len, 3), device=device, dtype=dtype)
|
||||
|
||||
#shift
|
||||
shift_f_size = 81 # Default value
|
||||
shift_f = False
|
||||
if shift_f:
|
||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(shift_f_size, shift_f_size + cond_f_len - 1,steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
else:
|
||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(0, cond_f_len - 1, steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
cond_img_ids[:, :, :, 1] = cond_img_ids[:, :, :, 1] + torch.linspace(h_len, h_len + cond_h_len - 1, steps=cond_h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
|
||||
cond_img_ids[:, :, :, 2] = cond_img_ids[:, :, :, 2] + torch.linspace(w_len, w_len + cond_w_len - 1, steps=cond_w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
|
||||
|
||||
# Combine original and conditional position ids
|
||||
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
cond_img_ids = repeat(cond_img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
combined_img_ids = torch.cat([img_ids, cond_img_ids], dim=1)
|
||||
|
||||
# Generate RoPE frequencies for the combined positions
|
||||
freqs = self.rope_embedder(combined_img_ids, ntk_alphas).movedim(1, 2)
|
||||
else:
|
||||
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
|
||||
return freqs
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -1611,7 +1958,13 @@ class WanModel(torch.nn.Module):
|
||||
mtv_motion_rotary_emb=None,
|
||||
mtv_freqs=None,
|
||||
mtv_strength=1.0,
|
||||
|
||||
s2v_audio_input=None,
|
||||
s2v_ref_latent=None,
|
||||
s2v_audio_scale=1.0,
|
||||
s2v_ref_motion=None,
|
||||
s2v_pose=None,
|
||||
s2v_motion_frames=[1, 0],
|
||||
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -1655,8 +2008,24 @@ class WanModel(torch.nn.Module):
|
||||
if isinstance(submodule, nn.Linear):
|
||||
if hasattr(submodule, 'step'):
|
||||
submodule.step = current_step
|
||||
|
||||
#s2v
|
||||
if self.model_type == 's2v' and s2v_audio_input is not None:
|
||||
if is_uncond:
|
||||
s2v_audio_input = s2v_audio_input * 0 # to match original code
|
||||
s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, s2v_motion_frames[0]), s2v_audio_input], dim=-1)
|
||||
|
||||
audio_emb_res = self.casual_audio_encoder(s2v_audio_input)
|
||||
if self.enable_adain:
|
||||
audio_emb_global, audio_emb = audio_emb_res
|
||||
self.audio_emb_global = audio_emb_global[:, s2v_motion_frames[1]:].clone()
|
||||
else:
|
||||
audio_emb = audio_emb_res
|
||||
merged_audio_emb = audio_emb[:, s2v_motion_frames[1]:, :]
|
||||
|
||||
# params
|
||||
device = self.patch_embedding.weight.device
|
||||
|
||||
if freqs is not None and freqs.device != device:
|
||||
freqs = freqs.to(device)
|
||||
|
||||
@@ -1697,23 +2066,30 @@ class WanModel(torch.nn.Module):
|
||||
for u in x
|
||||
]
|
||||
|
||||
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)
|
||||
|
||||
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)]
|
||||
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x])
|
||||
|
||||
grid_sizes = torch.stack([torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x])
|
||||
original_grid_sizes = grid_sizes.clone()
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
|
||||
x_len = x[0].shape[1]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.float32)
|
||||
assert seq_lens.max() <= seq_len
|
||||
|
||||
if self.trainable_cond_mask is not None:
|
||||
cond_mask_weight = self.trainable_cond_mask.weight.to(x[0]).unsqueeze(1).unsqueeze(1)
|
||||
|
||||
self.original_seq_len = x[0].shape[1]
|
||||
|
||||
if add_cond is not None:
|
||||
add_cond = self.add_conv_in(add_cond.to(self.add_conv_in.weight.dtype)).to(x[0].dtype)
|
||||
add_cond = add_cond.flatten(2).transpose(1, 2)
|
||||
x[0] = x[0] + self.add_proj(add_cond)
|
||||
if attn_cond is not None:
|
||||
F_cond, H_cond, W_cond = attn_cond.shape[2], attn_cond.shape[3], attn_cond.shape[4]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
attn_cond = self.attn_conv_in(attn_cond.to(self.attn_conv_in.weight.dtype)).to(x[0].dtype)
|
||||
attn_cond = attn_cond.flatten(2).transpose(1, 2)
|
||||
@@ -1727,24 +2103,34 @@ class WanModel(torch.nn.Module):
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
seq_len += fun_ref.size(1)
|
||||
F += 1
|
||||
x = [torch.concat([_fun_ref.unsqueeze(0), u], dim=1) for _fun_ref, u in zip(fun_ref, x)]
|
||||
x = [torch.cat([_fun_ref.unsqueeze(0), u], dim=1) for _fun_ref, u in zip(fun_ref, x)]
|
||||
|
||||
if phantom_ref is not None:
|
||||
phantom_ref_frames = phantom_ref.size(1)
|
||||
phantom_ref = self.original_patch_embedding(phantom_ref.unsqueeze(0).to(torch.float32)).flatten(2).transpose(1, 2).to(x[0].dtype)
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + phantom_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
phantom_ref_seq_len = phantom_ref.size(1)
|
||||
seq_len += phantom_ref_seq_len
|
||||
F += phantom_ref_frames
|
||||
x = [torch.concat([u, phantom_ref.unsqueeze(0)], dim=1) for phantom_ref, u in zip(phantom_ref, x)]
|
||||
end_ref_latent=None
|
||||
if s2v_ref_latent is not None:
|
||||
end_ref_latent = s2v_ref_latent.squeeze(0)
|
||||
elif phantom_ref is not None:
|
||||
end_ref_latent = phantom_ref
|
||||
F += end_ref_latent_frames
|
||||
if end_ref_latent is not None:
|
||||
end_ref_latent_frames = end_ref_latent.size(1)
|
||||
end_ref_latent = self.original_patch_embedding(end_ref_latent.unsqueeze(0).to(torch.float32)).to(x[0].dtype)
|
||||
end_ref_latent = end_ref_latent.flatten(2).transpose(1, 2)
|
||||
if cond_mask_weight is not None:
|
||||
end_ref_latent = end_ref_latent + cond_mask_weight[1]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + end_ref_latent_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
end_ref_latent_seq_len = end_ref_latent.size(1)
|
||||
seq_len += end_ref_latent_seq_len
|
||||
x = [torch.cat([u, end_ref_latent.unsqueeze(0)], dim=1) for end_ref_latent, u in zip(end_ref_latent, x)]
|
||||
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
|
||||
assert seq_lens.max() <= seq_len
|
||||
|
||||
x = torch.cat([
|
||||
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],
|
||||
dim=1) for u in x
|
||||
dim=1) for u in x
|
||||
])
|
||||
|
||||
if self.trainable_cond_mask is not None:
|
||||
x = x + cond_mask_weight[0]
|
||||
|
||||
# StandIn LoRA input
|
||||
x_ip = None
|
||||
freq_offset = 0
|
||||
@@ -1761,10 +2147,9 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
if freqs is None: #comfy rope
|
||||
current_shape = (F, H, W)
|
||||
|
||||
has_cond = attn_cond is not None
|
||||
f_len = ((F + (self.patch_size[0] // 2)) // self.patch_size[0])
|
||||
h_len = ((H + (self.patch_size[1] // 2)) // self.patch_size[1])
|
||||
w_len = ((W + (self.patch_size[2] // 2)) // self.patch_size[2])
|
||||
|
||||
if (self.cached_freqs is not None and
|
||||
self.cached_shape == current_shape and
|
||||
self.cached_cond == has_cond and
|
||||
@@ -1773,37 +2158,14 @@ class WanModel(torch.nn.Module):
|
||||
):
|
||||
freqs = self.cached_freqs
|
||||
else:
|
||||
img_ids = torch.zeros((f_len, h_len, w_len, 3), device=x.device, dtype=x.dtype)
|
||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(freq_offset, f_len + freq_offset - 1, steps=f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len + freq_offset - 1, steps=h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
|
||||
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len + freq_offset - 1, steps=w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
|
||||
|
||||
if attn_cond is not None:
|
||||
cond_f_len = ((F_cond + (self.patch_size[0] // 2)) // self.patch_size[0])
|
||||
cond_h_len = ((H_cond + (self.patch_size[1] // 2)) // self.patch_size[1])
|
||||
cond_w_len = ((W_cond + (self.patch_size[2] // 2)) // self.patch_size[2])
|
||||
cond_img_ids = torch.zeros((cond_f_len, cond_h_len, cond_w_len, 3), device=x.device, dtype=x.dtype)
|
||||
|
||||
#shift
|
||||
shift_f_size = 81 # Default value
|
||||
shift_f = False
|
||||
if shift_f:
|
||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(shift_f_size, shift_f_size + cond_f_len - 1,steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
else:
|
||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(0, cond_f_len - 1, steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
cond_img_ids[:, :, :, 1] = cond_img_ids[:, :, :, 1] + torch.linspace(h_len, h_len + cond_h_len - 1, steps=cond_h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
|
||||
cond_img_ids[:, :, :, 2] = cond_img_ids[:, :, :, 2] + torch.linspace(w_len, w_len + cond_w_len - 1, steps=cond_w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
|
||||
|
||||
# Combine original and conditional position ids
|
||||
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
cond_img_ids = repeat(cond_img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
combined_img_ids = torch.cat([img_ids, cond_img_ids], dim=1)
|
||||
|
||||
# Generate RoPE frequencies for the combined positions
|
||||
freqs = self.rope_embedder(combined_img_ids, ntk_alphas).movedim(1, 2)
|
||||
else:
|
||||
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
|
||||
freqs = self.rope_encode_comfy(F, H, W, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond=attn_cond, device=x.device, dtype=x.dtype)
|
||||
if s2v_ref_latent is not None:
|
||||
freqs_ref = self.rope_encode_comfy(
|
||||
s2v_ref_latent.shape[2],
|
||||
s2v_ref_latent.shape[3],
|
||||
s2v_ref_latent.shape[4],
|
||||
t_start=max(30, F + 9), device=x.device, dtype=x.dtype)
|
||||
freqs = torch.cat([freqs, freqs_ref], dim=1)
|
||||
|
||||
self.cached_freqs = freqs
|
||||
self.cached_shape = current_shape
|
||||
@@ -1814,6 +2176,8 @@ class WanModel(torch.nn.Module):
|
||||
# Stand-In RoPE frequencies
|
||||
if x_ip is not None:
|
||||
# Generate RoPE frequencies for x_ip
|
||||
h_len = (H + 1) // 2
|
||||
w_len = (W + 1) // 2
|
||||
ip_img_ids = torch.zeros((f_ip, h_ip, w_ip, 3), device=x.device, dtype=x.dtype)
|
||||
ip_img_ids[:, :, :, 0] = ip_img_ids[:, :, :, 0] + torch.linspace(0, f_ip - 1, steps=f_ip, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
ip_img_ids[:, :, :, 1] = ip_img_ids[:, :, :, 1] + torch.linspace(h_len + freq_offset, h_len + freq_offset + h_ip - 1, steps=h_ip, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
|
||||
@@ -1828,6 +2192,15 @@ class WanModel(torch.nn.Module):
|
||||
d = self.dim // self.num_heads
|
||||
self.cross_freqs = rope_params(100, d).to(device=x.device)
|
||||
|
||||
if s2v_ref_motion is not None:
|
||||
motion_encoded, freqs_motion = self.frame_packer(s2v_ref_motion, self)
|
||||
motion_encoded = motion_encoded + cond_mask_weight[2]
|
||||
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
|
||||
@@ -1835,9 +2208,23 @@ class WanModel(torch.nn.Module):
|
||||
else:
|
||||
expanded_timesteps = False
|
||||
|
||||
if self.zero_timestep:
|
||||
t = torch.cat([t, torch.zeros([1], dtype=t.dtype, device=t.device)])
|
||||
|
||||
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(x.dtype)) # b, dim
|
||||
e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim
|
||||
|
||||
#S2V zero timestep
|
||||
if self.zero_timestep:
|
||||
e = e[:-1]
|
||||
zero_e0 = e0[-1:]
|
||||
e0 = e0[:-1]
|
||||
e0 = torch.cat([
|
||||
e0.unsqueeze(2),
|
||||
zero_e0.unsqueeze(2).repeat(e0.size(0), 1, 1, 1)
|
||||
], dim=2)
|
||||
e0 = [e0, self.original_seq_len]
|
||||
|
||||
if x_ip is not None:
|
||||
timestep_ip = torch.zeros_like(t) # [B] with 0s
|
||||
t_ip = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, timestep_ip.flatten()).to(x.dtype)) # b, dim )
|
||||
@@ -2075,6 +2462,7 @@ class WanModel(torch.nn.Module):
|
||||
camera_embed=camera_embed,
|
||||
audio_proj=audio_proj,
|
||||
num_latent_frames = F,
|
||||
original_seq_len=self.original_seq_len,
|
||||
enhance_enabled=enhance_enabled,
|
||||
audio_scale=audio_scale,
|
||||
block_mask=self.block_mask,
|
||||
@@ -2171,6 +2559,8 @@ class WanModel(torch.nn.Module):
|
||||
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
|
||||
continue
|
||||
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.block_swap_debug:
|
||||
compute_end = time.perf_counter()
|
||||
compute_time = compute_end - compute_start
|
||||
@@ -2184,10 +2574,10 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
#uni3c controlnet
|
||||
if pdc_controlnet_states is not None and b < len(pdc_controlnet_states):
|
||||
x[:, :x_len] += pdc_controlnet_states[b].to(x) * pcd_data["controlnet_weight"]
|
||||
x[:, :self.original_seq_len] += pdc_controlnet_states[b].to(x) * pcd_data["controlnet_weight"]
|
||||
#controlnet
|
||||
if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])):
|
||||
x[:, :x_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"]
|
||||
x[:, :self.original_seq_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"]
|
||||
|
||||
if self.enable_teacache and (self.teacache_start_step <= current_step <= self.teacache_end_step) and pred_id is not None:
|
||||
self.teacache_state.update(
|
||||
@@ -2225,17 +2615,20 @@ class WanModel(torch.nn.Module):
|
||||
x = x[:, fun_ref_length:]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
if phantom_ref is not None:
|
||||
phantom_ref_length = phantom_ref.size(1)
|
||||
x = x[:, :-phantom_ref_length]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - phantom_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
if end_ref_latent is not None:
|
||||
end_ref_latent_length = end_ref_latent.size(1)
|
||||
x = x[:, :-end_ref_latent_length]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - end_ref_latent_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
if attn_cond is not None:
|
||||
x = x[:, :x_len]
|
||||
x = x[:, :self.original_seq_len]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
|
||||
x = x[:, :self.original_seq_len]
|
||||
|
||||
x = self.head(x, e.to(x.device))
|
||||
x = self.unpatchify(x, grid_sizes) # type: ignore[arg-type]
|
||||
x = self.unpatchify(x, original_grid_sizes) # type: ignore[arg-type]
|
||||
x = [u.float() for u in x]
|
||||
return (x, pred_id) if pred_id is not None else (x, None)
|
||||
|
||||
|
||||
@@ -0,0 +1,189 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import math
|
||||
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from transformers import Wav2Vec2ForCTC, Wav2Vec2Processor
|
||||
|
||||
|
||||
def get_sample_indices(original_fps,
|
||||
total_frames,
|
||||
target_fps,
|
||||
num_sample,
|
||||
fixed_start=None):
|
||||
required_duration = num_sample / target_fps
|
||||
required_origin_frames = int(np.ceil(required_duration * original_fps))
|
||||
if required_duration > total_frames / original_fps:
|
||||
raise ValueError("required_duration must be less than video length")
|
||||
|
||||
if not fixed_start is None and fixed_start >= 0:
|
||||
start_frame = fixed_start
|
||||
else:
|
||||
max_start = total_frames - required_origin_frames
|
||||
if max_start < 0:
|
||||
raise ValueError("video length is too short")
|
||||
start_frame = np.random.randint(0, max_start + 1)
|
||||
start_time = start_frame / original_fps
|
||||
|
||||
end_time = start_time + required_duration
|
||||
time_points = np.linspace(start_time, end_time, num_sample, endpoint=False)
|
||||
|
||||
frame_indices = np.round(np.array(time_points) * original_fps).astype(int)
|
||||
frame_indices = np.clip(frame_indices, 0, total_frames - 1)
|
||||
return frame_indices
|
||||
|
||||
|
||||
def linear_interpolation(features, input_fps, output_fps, output_len=None):
|
||||
"""
|
||||
features: shape=[1, T, 512]
|
||||
input_fps: fps for audio, f_a
|
||||
output_fps: fps for video, f_m
|
||||
output_len: video length
|
||||
"""
|
||||
features = features.transpose(1, 2) # [1, 512, T]
|
||||
seq_len = features.shape[2] / float(input_fps) # T/f_a
|
||||
if output_len is None:
|
||||
output_len = int(seq_len * output_fps) # f_m*T/f_a
|
||||
output_features = F.interpolate(
|
||||
features, size=output_len, align_corners=True,
|
||||
mode='linear') # [1, 512, output_len]
|
||||
return output_features.transpose(1, 2) # [1, output_len, 512]
|
||||
|
||||
|
||||
class AudioEncoder():
|
||||
|
||||
def __init__(self, device='cpu', model_id="facebook/wav2vec2-base-960h"):
|
||||
# load pretrained model
|
||||
self.processor = Wav2Vec2Processor.from_pretrained(model_id)
|
||||
self.model = Wav2Vec2ForCTC.from_pretrained(model_id)
|
||||
|
||||
self.model = self.model.to(device)
|
||||
|
||||
self.video_rate = 30
|
||||
|
||||
def extract_audio_feat(self,
|
||||
audio_path,
|
||||
return_all_layers=False,
|
||||
dtype=torch.float32):
|
||||
audio_input, sample_rate = librosa.load(audio_path, sr=16000)
|
||||
|
||||
input_values = self.processor(
|
||||
audio_input, sampling_rate=sample_rate,
|
||||
return_tensors="pt").input_values
|
||||
|
||||
# INFERENCE
|
||||
|
||||
# retrieve logits & take argmax
|
||||
res = self.model(
|
||||
input_values.to(self.model.device), output_hidden_states=True)
|
||||
if return_all_layers:
|
||||
feat = torch.cat(res.hidden_states)
|
||||
else:
|
||||
feat = res.hidden_states[-1]
|
||||
feat = linear_interpolation(
|
||||
feat, input_fps=50, output_fps=self.video_rate)
|
||||
|
||||
z = feat.to(dtype) # Encoding for the motion
|
||||
return z
|
||||
|
||||
def get_audio_embed_bucket(self,
|
||||
audio_embed,
|
||||
stride=2,
|
||||
batch_frames=12,
|
||||
m=2):
|
||||
num_layers, audio_frame_num, audio_dim = audio_embed.shape
|
||||
|
||||
if num_layers > 1:
|
||||
return_all_layers = True
|
||||
else:
|
||||
return_all_layers = False
|
||||
|
||||
min_batch_num = int(audio_frame_num / (batch_frames * stride)) + 1
|
||||
|
||||
bucket_num = min_batch_num * batch_frames
|
||||
batch_idx = [stride * i for i in range(bucket_num)]
|
||||
batch_audio_eb = []
|
||||
for bi in batch_idx:
|
||||
if bi < audio_frame_num:
|
||||
audio_sample_stride = 2
|
||||
chosen_idx = list(
|
||||
range(bi - m * audio_sample_stride,
|
||||
bi + (m + 1) * audio_sample_stride,
|
||||
audio_sample_stride))
|
||||
chosen_idx = [0 if c < 0 else c for c in chosen_idx]
|
||||
chosen_idx = [
|
||||
audio_frame_num - 1 if c >= audio_frame_num else c
|
||||
for c in chosen_idx
|
||||
]
|
||||
|
||||
if return_all_layers:
|
||||
frame_audio_embed = audio_embed[:, chosen_idx].flatten(
|
||||
start_dim=-2, end_dim=-1)
|
||||
else:
|
||||
frame_audio_embed = audio_embed[0][chosen_idx].flatten()
|
||||
else:
|
||||
frame_audio_embed = \
|
||||
torch.zeros([audio_dim * (2 * m + 1)], device=audio_embed.device) if not return_all_layers \
|
||||
else torch.zeros([num_layers, audio_dim * (2 * m + 1)], device=audio_embed.device)
|
||||
batch_audio_eb.append(frame_audio_embed)
|
||||
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb],
|
||||
dim=0)
|
||||
|
||||
return batch_audio_eb, min_batch_num
|
||||
|
||||
def get_audio_embed_bucket_fps(self,
|
||||
audio_embed,
|
||||
fps=16,
|
||||
batch_frames=81,
|
||||
m=0):
|
||||
num_layers, audio_frame_num, audio_dim = audio_embed.shape
|
||||
|
||||
if num_layers > 1:
|
||||
return_all_layers = True
|
||||
else:
|
||||
return_all_layers = False
|
||||
|
||||
scale = self.video_rate / fps
|
||||
|
||||
min_batch_num = int(audio_frame_num / (batch_frames * scale)) + 1
|
||||
|
||||
bucket_num = min_batch_num * batch_frames
|
||||
padd_audio_num = math.ceil(min_batch_num * batch_frames / fps *
|
||||
self.video_rate) - audio_frame_num
|
||||
batch_idx = get_sample_indices(
|
||||
original_fps=self.video_rate,
|
||||
total_frames=audio_frame_num + padd_audio_num,
|
||||
target_fps=fps,
|
||||
num_sample=bucket_num,
|
||||
fixed_start=0)
|
||||
batch_audio_eb = []
|
||||
audio_sample_stride = int(self.video_rate / fps)
|
||||
for bi in batch_idx:
|
||||
if bi < audio_frame_num:
|
||||
|
||||
chosen_idx = list(
|
||||
range(bi - m * audio_sample_stride,
|
||||
bi + (m + 1) * audio_sample_stride,
|
||||
audio_sample_stride))
|
||||
chosen_idx = [0 if c < 0 else c for c in chosen_idx]
|
||||
chosen_idx = [
|
||||
audio_frame_num - 1 if c >= audio_frame_num else c
|
||||
for c in chosen_idx
|
||||
]
|
||||
|
||||
if return_all_layers:
|
||||
frame_audio_embed = audio_embed[:, chosen_idx].flatten(
|
||||
start_dim=-2, end_dim=-1)
|
||||
else:
|
||||
frame_audio_embed = audio_embed[0][chosen_idx].flatten()
|
||||
else:
|
||||
frame_audio_embed = \
|
||||
torch.zeros([audio_dim * (2 * m + 1)], device=audio_embed.device) if not return_all_layers \
|
||||
else torch.zeros([num_layers, audio_dim * (2 * m + 1)], device=audio_embed.device)
|
||||
batch_audio_eb.append(frame_audio_embed)
|
||||
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb],
|
||||
dim=0)
|
||||
|
||||
return batch_audio_eb, min_batch_num
|
||||
@@ -0,0 +1,129 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
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 MotionEncoder_tc(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
in_dim: int,
|
||||
hidden_dim: int,
|
||||
num_heads=int,
|
||||
need_global=True,
|
||||
dtype=None,
|
||||
device=None):
|
||||
factory_kwargs = {"dtype": dtype, "device": device}
|
||||
super().__init__()
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.need_global = need_global
|
||||
self.conv1_local = CausalConv1d(
|
||||
in_dim, hidden_dim // 4 * num_heads, 3, stride=1)
|
||||
if need_global:
|
||||
self.conv1_global = CausalConv1d(
|
||||
in_dim, hidden_dim // 4, 3, stride=1)
|
||||
self.norm1 = nn.LayerNorm(
|
||||
hidden_dim // 4,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
**factory_kwargs)
|
||||
self.act = nn.SiLU()
|
||||
self.conv2 = CausalConv1d(hidden_dim // 4, hidden_dim // 2, 3, stride=2)
|
||||
self.conv3 = CausalConv1d(hidden_dim // 2, hidden_dim, 3, stride=2)
|
||||
|
||||
if need_global:
|
||||
self.final_linear = nn.Linear(hidden_dim, hidden_dim,
|
||||
**factory_kwargs)
|
||||
|
||||
self.norm1 = nn.LayerNorm(
|
||||
hidden_dim // 4,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
**factory_kwargs)
|
||||
|
||||
self.norm2 = nn.LayerNorm(
|
||||
hidden_dim // 2,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
**factory_kwargs)
|
||||
|
||||
self.norm3 = nn.LayerNorm(
|
||||
hidden_dim, elementwise_affine=False, eps=1e-6, **factory_kwargs)
|
||||
|
||||
self.padding_tokens = nn.Parameter(torch.zeros(1, 1, 1, hidden_dim))
|
||||
|
||||
def forward(self, x):
|
||||
x = rearrange(x, 'b t c -> b c t')
|
||||
x_ori = x.clone()
|
||||
b, c, t = x.shape
|
||||
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 = rearrange(x, '(b n) t c -> b t n c', b=b)
|
||||
padding = self.padding_tokens.repeat(b, x.shape[1], 1, 1)
|
||||
x = torch.cat([x, padding], dim=-2)
|
||||
x_local = x.clone()
|
||||
|
||||
if not self.need_global:
|
||||
return x_local
|
||||
|
||||
x = self.conv1_global(x_ori)
|
||||
x = rearrange(x, 'b c t -> b t c')
|
||||
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.final_linear(x)
|
||||
x = rearrange(x, '(b n) t c -> b t n c', b=b)
|
||||
|
||||
return x, x_local
|
||||
@@ -22,7 +22,7 @@ scheduler_list = [
|
||||
"multitalk"
|
||||
]
|
||||
|
||||
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None):
|
||||
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None, log_timesteps=False):
|
||||
timesteps = None
|
||||
if 'unipc' in scheduler:
|
||||
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
|
||||
@@ -111,7 +111,8 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
start_idx = 0
|
||||
end_idx = len(timesteps) - 1
|
||||
|
||||
log.info(f"Total timesteps: {timesteps}")
|
||||
if log_timesteps:
|
||||
log.info(f"Total timesteps: {timesteps}")
|
||||
|
||||
if isinstance(start_step, float):
|
||||
idxs = (sample_scheduler.sigmas <= start_step).nonzero(as_tuple=True)[0]
|
||||
@@ -134,8 +135,8 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
sample_scheduler.full_sigmas = sample_scheduler.sigmas.clone()
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer
|
||||
|
||||
|
||||
log.info(f"Using timesteps: {timesteps}")
|
||||
if log_timesteps:
|
||||
log.info(f"Using timesteps: {timesteps}")
|
||||
|
||||
if hasattr(sample_scheduler, 'timesteps'):
|
||||
sample_scheduler.timesteps = timesteps
|
||||
|
||||
Reference in New Issue
Block a user