diff --git a/nodes_sampler.py b/nodes_sampler.py index 7e4ff44..f12233c 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -11,17 +11,15 @@ from .custom_linear import remove_lora_from_module, set_lora_params, _replace_li from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps, scheduler_list from .gguf.gguf import set_lora_params_gguf from .multitalk.multitalk import timestep_transform, add_noise -from .utils import(log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, optimized_scale, setup_radial_attention, - compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device, get_raag_guidance, temporal_score_rescaling) +from .utils import(log, print_memory, apply_lora, fourier_filter, optimized_scale, setup_radial_attention, + compile_model, dict_to_device, tangential_projection, get_raag_guidance, temporal_score_rescaling) from .cache_methods.cache_methods import cache_report from .nodes_model_loading import load_weights from .enhance_a_video.globals import set_enhance_weight, set_num_frames from contextlib import nullcontext -from einops import rearrange from comfy import model_management as mm from comfy.utils import ProgressBar, load_torch_file -from comfy.clip_vision import clip_preprocess, ClipVisionModel from comfy.cli_args import args, LatentPreviewMethod script_directory = os.path.dirname(os.path.abspath(__file__)) @@ -46,7 +44,7 @@ class MetaParameter(torch.nn.Parameter): self.quant_type = quant_type return self -def offload_transformer(transformer): +def offload_transformer(transformer): transformer.teacache_state.clear_all() transformer.magcache_state.clear_all() transformer.easycache_state.clear_all() @@ -199,7 +197,7 @@ class WanVideoSampler: transformer = _replace_linear(transformer, dtype, patcher.model["sd"], compile_args=model["compile_args"]) transformer.patched_linear = True if patcher.model["sd"] is not None and gguf_reader is None: - load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, + load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args, compile_args=model["compile_args"]) if gguf_reader is not None: #handle GGUF @@ -344,7 +342,7 @@ class WanVideoSampler: dtype=torch.float32, generator=seed_g, device=torch.device("cpu")) - + seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1]) control_embeds = image_embeds.get("control_embeds", None) @@ -414,7 +412,7 @@ class WanVideoSampler: dtype=torch.float32, device=torch.device("cpu"), generator=seed_g) - + seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1]) recammaster = image_embeds.get("recammaster", None) @@ -425,7 +423,7 @@ class WanVideoSampler: log.info(f"RecamMaster camera embed shape: {camera_embed.shape}") log.info(f"RecamMaster source video shape: {recam_latents.shape}") seq_len *= 2 - + if image_embeds.get("mocha_embeds", None) is not None: mocha_embeds = image_embeds.get("mocha_embeds", None) mocha_num_refs = image_embeds.get("mocha_num_refs", 0) @@ -1012,7 +1010,7 @@ class WanVideoSampler: if experimental_args is not None: video_attention_split_steps = experimental_args.get("video_attention_split_steps", []) if video_attention_split_steps: - transformer.video_attention_split_steps = [int(x.strip()) for x in video_attention_split_steps.split(",")] + transformer.video_attention_split_steps = [int(x.strip()) for x in video_attention_split_steps.split(",")] use_zero_init = experimental_args.get("use_zero_init", True) use_cfg_zero_star = experimental_args.get("cfg_zero_star", False) @@ -1055,7 +1053,7 @@ class WanVideoSampler: from .mocha.nodes import rope_params_mocha log.info(f"Using Mocha RoPE") rope_function = 'mocha' - + freqs = torch.cat([ rope_params_mocha(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index, start=-1), rope_params_mocha(1024, 2 * (d // 6), start=-1), @@ -1105,7 +1103,7 @@ class WanVideoSampler: log.info(f"Lynx ref latent shape: {lynx_ref_latent.shape}") log.info("Extracting Lynx ref cond buffer...") if transformer.in_dim == 36: - mask_latents = torch.tile(torch.zeros_like(lynx_ref_latent[:1]), [4, 1, 1, 1]) + mask_latents = torch.tile(torch.zeros_like(lynx_ref_latent[:1]), [4, 1, 1, 1]) empty_image_cond = torch.cat([mask_latents, torch.zeros_like(lynx_ref_latent)], dim=0).to(device) lynx_ref_input = torch.cat([lynx_ref_latent, empty_image_cond], dim=0) else: @@ -1133,7 +1131,7 @@ class WanVideoSampler: is_uncond=True ) log.info(f"Extracted {len(lynx_ref_buffer_uncond)} uncond ref buffers") - + if lynx_embeds.get("ip_x", None) is not None: lynx_embeds["ip_x"] = lynx_embeds["ip_x"].to(device, dtype) lynx_embeds["ip_x_uncond"] = lynx_embeds["ip_x_uncond"].to(device, dtype) @@ -1262,7 +1260,7 @@ class WanVideoSampler: if recammaster is not None: z = torch.cat([z, recam_latents.to(z)], dim=1) - + if mocha_embeds is not None: if context_window is not None and mocha_embeds.shape[2] != context_frames: latent_frames = len(context_window) @@ -1272,7 +1270,7 @@ class WanVideoSampler: partial_latents = mocha_embeds[:, context_window] # windowed latents mask_frame = mocha_embeds[:, latent_end:mask_end] # single mask frame ref_frames = mocha_embeds[:, -mocha_num_refs:] # reference frames - + partial_mocha_embeds = torch.cat([partial_latents, mask_frame, ref_frames], dim=1) z = torch.cat([z, partial_mocha_embeds.to(z)], dim=1) else: @@ -1691,6 +1689,7 @@ class WanVideoSampler: # Main sampling loop with FreeInit iterations iterations = freeinit_args.get("freeinit_num_iters", 3) if freeinit_args is not None else 1 current_latent = latent + initial_noise_saved = None for iter_idx in range(iterations): @@ -2370,7 +2369,6 @@ class WanVideoSampler: vae.to(device) # Pad original_images if needed num_frames = original_images.shape[2] - required_frames = audio_end_idx - audio_start_idx if audio_end_idx > num_frames: pad_len = audio_end_idx - num_frames last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1) @@ -2536,9 +2534,9 @@ class WanVideoSampler: source_frame = len(audio_embedding[human_inx]) source_frames.append(source_frame) if audio_end_idx >= len(audio_embedding[human_inx]): - print(f"Audio embedding for subject {human_inx} not long enough: {len(audio_embedding[human_inx])}, need {audio_end_idx}, padding...") + log.warning(f"Audio embedding for subject {human_inx} not long enough: {len(audio_embedding[human_inx])}, need {audio_end_idx}, padding...") miss_length = audio_end_idx - len(audio_embedding[human_inx]) + 3 - print(f"Padding length: {miss_length}") + log.warning(f"Padding length: {miss_length}") if encoded_silence is not None: add_audio_emb = encoded_silence[-1*miss_length:] else: @@ -2817,7 +2815,7 @@ class WanVideoSampler: vae.to(device) pose_image_slice = pose_images_in[:, start:end].to(device) pose_input_slice = vae.encode([pose_image_slice], device,tiled=tiled_vae, pbar=False).to(dtype) - + vae.to(offload_device) if wananim_face_pixels is None and wananim_ref_masks is not None: @@ -3006,7 +3004,7 @@ class WanVideoSampler: #region normal inference else: noise_pred, noise_pred_ovi, self.cache_state = predict_with_cfg( - latent_model_input, + 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, multitalk_audio_embeds=multitalk_audio_embeds, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input, @@ -3075,7 +3073,7 @@ class WanVideoSampler: **scheduler_step_args)[0].squeeze(0) latent_backwards = torch.flip(latent_backwards, dims=[1]) latent = latent * 0.5 + latent_backwards * 0.5 - + if latent_ovi is not None: latent_ovi = sample_scheduler_ovi.step(noise_pred_ovi.unsqueeze(0), t, latent_ovi.to(device).unsqueeze(0), **scheduler_step_args)[0].squeeze(0) @@ -3097,7 +3095,7 @@ class WanVideoSampler: ) mask = masks[idx].to(latent) latent = image_latent * mask + latent * (1-mask) - + # TTM if ttm_reference_latents is not None and (idx + ttm_start_step) < ttm_end_step: if idx + ttm_start_step + 1 < len(sample_scheduler.all_timesteps):