Code cleanup
This commit is contained in:
+20
-22
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user