Code cleanup

This commit is contained in:
kijai
2025-12-02 14:49:42 +02:00
parent e60eb995ee
commit aebeeb9160
+20 -22
View File
@@ -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):