Implement Framepack long geneneration method
RoPE handling is comfyanon's code
This commit is contained in:
@@ -1823,6 +1823,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)
|
||||
|
||||
@@ -2223,14 +2225,19 @@ class WanVideoSampler:
|
||||
mtv_freqs = mtv_freqs.to(device, dtype)
|
||||
|
||||
#region S2V
|
||||
s2v_audio_input = s2v_ref_latent = None
|
||||
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"]].to(device, dtype)
|
||||
#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:
|
||||
@@ -2241,10 +2248,10 @@ class WanVideoSampler:
|
||||
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 = image_embeds.get("vae", None)
|
||||
framepack = False
|
||||
vae = s2v_audio_embeds.get("vae", None)
|
||||
|
||||
# vid2vid
|
||||
noise_mask=original_image=None
|
||||
@@ -2525,7 +2532,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, s2v_audio_input=None, s2v_ref_motion=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)
|
||||
@@ -2670,7 +2677,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
|
||||
@@ -2707,7 +2718,8 @@ class WanVideoSampler:
|
||||
"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_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
|
||||
@@ -2862,7 +2874,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")
|
||||
|
||||
@@ -3229,6 +3241,10 @@ class WanVideoSampler:
|
||||
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)
|
||||
@@ -3246,7 +3262,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, s2v_audio_input=partial_s2v_audio_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)
|
||||
|
||||
if cache_args is not None:
|
||||
self.window_tracker.cache_states[window_id] = new_teacache
|
||||
@@ -3671,67 +3687,124 @@ 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
|
||||
motion_frames = 5
|
||||
infer_frames = image_embeds["num_frames"]
|
||||
#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
|
||||
|
||||
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):
|
||||
if ref_motion_image is not None:
|
||||
if ref_motion_image.shape[0] > 73:
|
||||
ref_motion_image = ref_motion_image[-73:]
|
||||
|
||||
if ref_motion_image.shape[0] < 73:
|
||||
ref = torch.ones([73, ref_motion_image.shape[1], ref_motion_image.shape[2], 3]) * 0.5
|
||||
ref[-ref_motion_image.shape[0]:] = ref_motion_image
|
||||
ref_motion_image = ref
|
||||
|
||||
vae.to(device)
|
||||
ref_motion = vae.encode(ref_motion_image[:, :, :, :3], device=device, pbar=False)[0].to(dtype)
|
||||
ref_motion = vae.encode(ref_motion_image.to(vae.dtype), device=device, pbar=False).to(dtype)[0]
|
||||
vae.to(offload_device)
|
||||
|
||||
left_idx = r * infer_frames
|
||||
right_idx = r * infer_frames + infer_frames
|
||||
#cond_latents = COND[r] if pose_video else COND[0] * 0
|
||||
#cond_latents = cond_latents.to(dtype=self.param_dtype, device=self.device)
|
||||
s2v_audio_input = s2v_audio_embeds[..., left_idx:right_idx]
|
||||
input_motion_latents = ref_motion.clone()
|
||||
|
||||
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, s2v_ref_motion=input_motion_latents)
|
||||
|
||||
latent = sample_scheduler.step(
|
||||
noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0),
|
||||
**scheduler_step_args)[0].squeeze(0)
|
||||
|
||||
latents = torch.stack(latent)
|
||||
#if not (drop_first_motion and r == 0):
|
||||
# decode_latents = torch.cat([motion_latents, latents], dim=2)
|
||||
#else:
|
||||
decode_latents = torch.cat([s2v_ref_latent, latents], dim=2)
|
||||
image = torch.stack(vae.decode(decode_latents), device=device)
|
||||
image = image[:, :, -(infer_frames):]
|
||||
#if (drop_first_motion and r == 0):
|
||||
# image = image[:, :, 3:]
|
||||
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)
|
||||
|
||||
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(vae.device, vae.dtype)
|
||||
|
||||
if ref_motion_image is not None:
|
||||
input_motion_latents = ref_motion.clone().unsqueeze(0)
|
||||
else:
|
||||
input_motion_latents = 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
|
||||
|
||||
|
||||
vae.to(device)
|
||||
ref_motion_image = torch.stack(vae.encode(videos_last_frames, device=device, pbar=False)[0])
|
||||
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]
|
||||
image = image.unsqueeze(0)[:, :, -infer_frames:]
|
||||
if r == 0:
|
||||
image = image[:, :, 3:]
|
||||
|
||||
framepack_out.append(image.cpu())
|
||||
|
||||
gen_video_samples = torch.cat(framepack_out, dim=1)
|
||||
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)
|
||||
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"]:
|
||||
@@ -3741,7 +3814,7 @@ class WanVideoSampler:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return {"video": gen_video_samples.permute(1, 2, 3, 0)},
|
||||
return {"video": gen_video_samples},
|
||||
|
||||
#region normal inference
|
||||
else:
|
||||
|
||||
@@ -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", "cond_encoder"}
|
||||
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
|
||||
@@ -878,6 +879,7 @@ def patch_stand_in_lora(transformer, lora_sd, transformer_load_device, base_dtyp
|
||||
|
||||
def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
|
||||
unianimate_sd = None
|
||||
control_lora=False
|
||||
#spacepxl's control LoRA patch
|
||||
for l in lora:
|
||||
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
|
||||
@@ -904,7 +906,6 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
|
||||
# Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys
|
||||
#if not any('img' in k for k in sd.keys()):
|
||||
# lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k}
|
||||
control_lora=False
|
||||
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
|
||||
control_lora = True
|
||||
#stand-in LoRA patch
|
||||
|
||||
+33
-19
@@ -46,47 +46,55 @@ def linear_interpolation(features, input_fps, output_fps, output_len=None):
|
||||
mode='linear') # [1, 512, output_len]
|
||||
return output_features.transpose(1, 2) # [1, output_len, 512]
|
||||
|
||||
class WanVideoAddAudioEmbeds:
|
||||
class WanVideoAddS2VEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"frames": ("INT", {"default": 81, "min": 1, "max": 100000, "step": 1, "tooltip": "Number of frames to process"}),
|
||||
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for audio embeddings"})
|
||||
"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",)
|
||||
"pose_latent": ("LATENT",),
|
||||
"vae": ("WANVAE",),
|
||||
"enable_framepack": ("BOOLEAN", {"default": False, "tooltip": "Enable Framepack sampling loop, not compatible with context windows"})
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "INT",)
|
||||
RETURN_NAMES = ("image_embeds", "audio_frame_count")
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, frames, audio_encoder_output=None, audio_scale=1.0, ref_latent=None, pose_latent=None):
|
||||
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):
|
||||
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", audio_feat.shape)
|
||||
input_fps = 50
|
||||
output_fps = 30
|
||||
bucket_fps = 16
|
||||
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=frames-1
|
||||
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:
|
||||
@@ -94,6 +102,8 @@ class WanVideoAddAudioEmbeds:
|
||||
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 = {
|
||||
@@ -101,11 +111,16 @@ class WanVideoAddAudioEmbeds:
|
||||
"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
|
||||
"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,)
|
||||
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
|
||||
@@ -120,8 +135,7 @@ class WanVideoAddAudioEmbeds:
|
||||
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
|
||||
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,
|
||||
@@ -161,9 +175,9 @@ class WanVideoAddAudioEmbeds:
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddAudioEmbeds": WanVideoAddAudioEmbeds,
|
||||
"WanVideoAddS2VEmbeds": WanVideoAddS2VEmbeds,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddAudioEmbeds": "WanVideo Add Audio Embeds",
|
||||
"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
+23
-26
@@ -40,11 +40,11 @@ class FramePackMotioner(nn.Module):
|
||||
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
|
||||
dtype=None, device=None):
|
||||
):
|
||||
super().__init__()
|
||||
self.proj = nn.Conv3d(16, inner_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), dtype=dtype, device=device)
|
||||
self.proj_2x = nn.Conv3d(16, inner_dim, kernel_size=(2, 4, 4), stride=(2, 4, 4), dtype=dtype, device=device)
|
||||
self.proj_4x = nn.Conv3d(16, inner_dim, kernel_size=(4, 8, 8), stride=(4, 8, 8), dtype=dtype, device=device)
|
||||
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
|
||||
@@ -79,9 +79,9 @@ class FramePackMotioner(nn.Module):
|
||||
|
||||
motion_lat = torch.cat([clean_latents_post, clean_latents_2x, clean_latents_4x], dim=1)
|
||||
|
||||
rope_post = rope_embedder.rope_encode(1, lat_height, lat_width, t_start=-1, device=motion_latents.device, dtype=motion_latents.dtype)
|
||||
rope_2x = rope_embedder.rope_encode(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(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_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
|
||||
@@ -1714,15 +1714,13 @@ class WanModel(torch.nn.Module):
|
||||
# WanLayerNorm(motioner_dim),
|
||||
# zero_module(nn.Linear(motioner_dim, self.dim)))
|
||||
|
||||
self.enable_framepack = enable_framepack
|
||||
enable_framepack = True
|
||||
if enable_framepack:
|
||||
self.frame_packer = FramePackMotioner(
|
||||
inner_dim=self.dim,
|
||||
num_heads=self.num_heads,
|
||||
zip_frame_buckets=[1, 2, 16],
|
||||
drop_mode='padd',
|
||||
device=self.main_device,
|
||||
dtype=self.dtype)
|
||||
drop_mode='padd')
|
||||
|
||||
@staticmethod
|
||||
def _prepare_blockwise_causal_attn_mask(
|
||||
@@ -2139,7 +2137,8 @@ class WanModel(torch.nn.Module):
|
||||
s2v_ref_latent=None,
|
||||
s2v_audio_scale=1.0,
|
||||
s2v_ref_motion=None,
|
||||
s2v_pose=None
|
||||
s2v_pose=None,
|
||||
s2v_motion_frames=[1, 0],
|
||||
|
||||
):
|
||||
r"""
|
||||
@@ -2189,17 +2188,15 @@ class WanModel(torch.nn.Module):
|
||||
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
|
||||
#motion_frames=[17, 5]
|
||||
motion_frames=[1, 0]
|
||||
s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, motion_frames[0]), s2v_audio_input], dim=-1)
|
||||
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[:, motion_frames[1]:].clone()
|
||||
self.audio_emb_global = audio_emb_global[:, s2v_motion_frames[1]:].clone()
|
||||
else:
|
||||
audio_emb = audio_emb_res
|
||||
merged_audio_emb = audio_emb[:, motion_frames[1]:, :]
|
||||
merged_audio_emb = audio_emb[:, s2v_motion_frames[1]:, :]
|
||||
|
||||
# params
|
||||
device = self.patch_embedding.weight.device
|
||||
@@ -2245,16 +2242,14 @@ class WanModel(torch.nn.Module):
|
||||
]
|
||||
|
||||
if s2v_pose is not None:
|
||||
print("s2v_pose.shape:", s2v_pose.shape)
|
||||
print("x[0].shape:", x[0].shape)
|
||||
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])
|
||||
original_grid_sizes = grid_sizes.clone()
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.float32)
|
||||
@@ -2302,7 +2297,7 @@ class WanModel(torch.nn.Module):
|
||||
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)]
|
||||
|
||||
grid_sizes = grid_sizes
|
||||
|
||||
x = torch.cat([
|
||||
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],
|
||||
dim=1) for u in x
|
||||
@@ -2344,7 +2339,7 @@ class WanModel(torch.nn.Module):
|
||||
s2v_ref_latent.shape[2],
|
||||
s2v_ref_latent.shape[3],
|
||||
s2v_ref_latent.shape[4],
|
||||
t_start=30, device=x.device, dtype=x.dtype)
|
||||
t_start=max(30, F + 9), device=x.device, dtype=x.dtype)
|
||||
freqs = torch.cat([freqs, freqs_ref], dim=1)
|
||||
|
||||
self.cached_freqs = freqs
|
||||
@@ -2746,10 +2741,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(
|
||||
@@ -2796,9 +2791,11 @@ class WanModel(torch.nn.Module):
|
||||
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 = 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)
|
||||
|
||||
|
||||
@@ -133,9 +133,6 @@ 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"timesteps: {timesteps}")
|
||||
|
||||
if hasattr(sample_scheduler, 'timesteps'):
|
||||
sample_scheduler.timesteps = timesteps
|
||||
|
||||
|
||||
Reference in New Issue
Block a user