Cleanup: Remove flowedit code, restructure some other code

Due to lack of use and code maintainability
This commit is contained in:
kijai
2025-12-26 12:51:22 +02:00
parent 74f337e06c
commit 20942b8fd9
5 changed files with 673 additions and 819 deletions
+483
View File
@@ -0,0 +1,483 @@
import torch
import os
import gc
from PIL import Image
import numpy as np
from ..latent_preview import prepare_callback
from ..wanvideo.schedulers import get_scheduler
from .multitalk import timestep_transform, add_noise
from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap
from comfy.utils import load_torch_file
from ..nodes_model_loading import load_weights
from ..HuMo.nodes import get_audio_emb_window
import comfy.model_management as mm
from tqdm import tqdm
import copy
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
vae_upscale_factor = 16
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
def multitalk_loop(self, **kwargs):
# Unpack kwargs into local variables
(latent, total_steps, steps, start_step, end_step, shift, cfg, denoise_strength,
sigmas, weight_dtype, transformer, patcher, block_swap_args, model, vae, dtype,
scheduler, scheduler_step_args, text_embeds, image_embeds, multitalk_embeds,
multitalk_audio_embeds, unianim_data, dwpose_data, unianimate_poses, uni3c_embeds,
humo_image_cond, humo_image_cond_neg, humo_audio, humo_reference_count,
add_noise_to_samples, audio_stride, use_tsr, tsr_k, tsr_sigma, fantasy_portrait_input,
noise, timesteps, force_offload, add_cond, control_latents, audio_proj,
control_camera_latents, samples, masks, seed_g, gguf_reader, predict_func
) = (kwargs.get(k) for k in (
'latent', 'total_steps', 'steps', 'start_step', 'end_step', 'shift', 'cfg',
'denoise_strength', 'sigmas', 'weight_dtype', 'transformer', 'patcher',
'block_swap_args', 'model', 'vae', 'dtype', 'scheduler', 'scheduler_step_args',
'text_embeds', 'image_embeds', 'multitalk_embeds', 'multitalk_audio_embeds',
'unianim_data', 'dwpose_data', 'unianimate_poses', 'uni3c_embeds',
'humo_image_cond', 'humo_image_cond_neg', 'humo_audio', 'humo_reference_count',
'add_noise_to_samples', 'audio_stride', 'use_tsr', 'tsr_k', 'tsr_sigma',
'fantasy_portrait_input', 'noise', 'timesteps', 'force_offload', 'add_cond',
'control_latents', 'audio_proj', 'control_camera_latents', 'samples', 'masks',
'seed_g', 'gguf_reader', 'predict_with_cfg'
))
mode = image_embeds.get("multitalk_mode", "multitalk")
if mode == "auto":
mode = transformer.multitalk_model_type.lower()
log.info(f"Multitalk mode: {mode}")
cond_frame = None
offload = image_embeds.get("force_offload", False)
offloaded = False
tiled_vae = image_embeds.get("tiled_vae", False)
frame_num = clip_length = image_embeds.get("frame_window_size", 81)
clip_embeds = image_embeds.get("clip_context", None)
if clip_embeds is not None:
clip_embeds = clip_embeds.to(dtype)
colormatch = image_embeds.get("colormatch", "disabled")
motion_frame = image_embeds.get("motion_frame", 25)
target_w = image_embeds.get("target_w", None)
target_h = image_embeds.get("target_h", None)
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
if original_images is None:
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
output_path = image_embeds.get("output_path", "")
img_counter = 0
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
face_scale = 0.1
x_min, x_max = int(target_h * face_scale), int(target_h * (1 - face_scale))
lefty_min, lefty_max = int((target_w//2) * face_scale), int((target_w//2) * (1 - face_scale))
righty_min, righty_max = int((target_w//2) * face_scale + (target_w//2)), int((target_w//2) * (1 - face_scale) + (target_w//2))
human_mask1, human_mask2 = (torch.zeros([target_h, target_w]) for _ in range(2))
human_mask1[x_min:x_max, lefty_min:lefty_max] = 1
human_mask2[x_min:x_max, righty_min:righty_max] = 1
background_mask = torch.where((human_mask1 + human_mask2) > 0, torch.tensor(0), torch.tensor(1))
human_masks = [human_mask1, human_mask2, background_mask]
ref_target_masks = torch.stack(human_masks, dim=0)
multitalk_embeds['ref_target_masks'] = ref_target_masks
gen_video_list = []
is_first_clip = True
arrive_last_frame = False
cur_motion_frames_num = 1
audio_start_idx = iteration_count = step_iteration_count = 0
audio_end_idx = (audio_start_idx + clip_length) * audio_stride
indices = (torch.arange(4 + 1) - 2) * 1
current_condframe_index = 0
audio_embedding = multitalk_audio_embeds
human_num = len(audio_embedding)
audio_embs = None
cond_frame = None
uni3c_data = None
if uni3c_embeds is not None:
transformer.controlnet = uni3c_embeds["controlnet"]
uni3c_data = uni3c_embeds.copy()
encoded_silence = None
try:
silence_path = os.path.join(script_directory, "encoded_silence.safetensors")
encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype)
except:
log.warning("No encoded silence file found, padding with end of audio embedding instead.")
total_frames = len(audio_embedding[0])
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
callback = prepare_callback(patcher, estimated_iterations)
if frame_num >= total_frames:
arrive_last_frame = True
estimated_iterations = 1
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
while True: # start video generation iteratively
self.cache_state = [None, None]
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
if mode == "infinitetalk":
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None
if multitalk_embeds is not None:
audio_embs = []
# split audio with window size
for human_idx in range(human_num):
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0]-1)
audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device)
audio_embs.append(audio_emb)
audio_embs = torch.concat(audio_embs, dim=0).to(dtype)
h, w = (cond_image.shape[-2], cond_image.shape[-1]) if cond_image is not None else (target_h, target_w)
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
latent_frame_num = (frame_num - 1) // 4 + 1
noise = torch.randn(
16, latent_frame_num,
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
# Calculate the correct latent slice based on current iteration
if is_first_clip:
latent_start_idx = 0
latent_end_idx = noise.shape[1]
else:
new_frames_per_iteration = frame_num - motion_frame
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
latent_start_idx = iteration_count * new_latent_frames_per_iteration
latent_end_idx = latent_start_idx + noise.shape[1]
if samples is not None:
noise_mask = samples.get("noise_mask", None)
input_samples = samples["samples"]
if input_samples is not None:
input_samples = input_samples.squeeze(0).to(noise)
# Check if we have enough frames in input_samples
if latent_end_idx > input_samples.shape[1]:
# We need more frames than available - pad the input_samples at the end
pad_length = latent_end_idx - input_samples.shape[1]
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
input_samples = torch.cat([input_samples, last_frame], dim=1)
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
if noise_mask is not None:
original_image = input_samples.to(device)
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
if add_noise_to_samples:
latent_timestep = timesteps[0]
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
else:
noise = input_samples
# diff diff prep
if noise_mask is not None:
if len(noise_mask.shape) == 4:
noise_mask = noise_mask.squeeze(1)
if audio_end_idx > noise_mask.shape[0]:
noise_mask = noise_mask.repeat(audio_end_idx // noise_mask.shape[0], 1, 1)
noise_mask = noise_mask[audio_start_idx:audio_end_idx]
noise_mask = torch.nn.functional.interpolate(
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
mode='trilinear',
align_corners=False
).repeat(1, noise.shape[0], 1, 1, 1)
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device)
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
# zero padding and vae encode for img cond
if cond_image is not None or cond_frame is not None:
cond_ = cond_image if (is_first_clip or humo_image_cond is None) else cond_frame
cond_frame_num = cond_.shape[2]
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num, target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.concat([cond_.to(device, vae.dtype), video_frames], dim=2)
# encode
vae.to(device)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
if mode == "multitalk":
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
else:
cond_ = cond_image if is_first_clip else cond_frame
latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
vae.to(offload_device)
#motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
msk[:, :1] = 1
y = torch.cat([msk, y]) # 4+C T H W
mm.soft_empty_cache()
else:
y = None
latent_motion_frames = noise[:, :1]
partial_humo_cond_input = partial_humo_cond_neg_input = partial_humo_audio = partial_humo_audio_neg = None
if humo_image_cond is not None:
partial_humo_cond_input = humo_image_cond[:, :latent_frame_num]
partial_humo_cond_neg_input = humo_image_cond_neg[:, :latent_frame_num]
if y is not None:
partial_humo_cond_input[:, :1] = y[:, :1]
if humo_reference_count > 0:
partial_humo_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:]
partial_humo_cond_neg_input[:, -humo_reference_count:] = humo_image_cond_neg[:, -humo_reference_count:]
if humo_audio is not None:
if is_first_clip:
audio_embs = None
partial_humo_audio, _ = get_audio_emb_window(humo_audio, frame_num, frame0_idx=audio_start_idx)
#zero_audio_pad = torch.zeros(humo_reference_count, *partial_humo_audio.shape[1:], device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
partial_humo_audio[-humo_reference_count:] = 0
partial_humo_audio_neg = torch.zeros_like(partial_humo_audio, device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
if scheduler == "multitalk":
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
timesteps.append(0.)
timesteps = [torch.tensor([t], device=device) for t in timesteps]
timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps]
else:
if isinstance(scheduler, dict):
sample_scheduler = copy.deepcopy(scheduler["sample_scheduler"])
timesteps = scheduler["timesteps"]
else:
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, denoise_strength, sigmas=sigmas)
timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)]
# sample videos
latent = noise
# injecting motion frames
if not is_first_clip and mode == "multitalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
latent[:, :add_latent.shape[1]] = add_latent
if offloaded:
# 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)
elif 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)
#blockswap init
init_blockswap(transformer, block_swap_args, model)
# Use the appropriate prompt for this section
if len(text_embeds["prompt_embeds"]) > 1:
prompt_index = min(iteration_count, len(text_embeds["prompt_embeds"]) - 1)
positive = [text_embeds["prompt_embeds"][prompt_index]]
log.info(f"Using prompt index: {prompt_index}")
else:
positive = text_embeds["prompt_embeds"]
# uni3c slices
if uni3c_embeds is not None:
vae.to(device)
# Pad original_images if needed
num_frames = original_images.shape[2]
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)
padded_images = torch.cat([original_images, last_frame], dim=2)
else:
padded_images = original_images
render_latent = vae.encode(
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
device=device, tiled=tiled_vae
).to(dtype)
vae.to(offload_device)
uni3c_data['render_latent'] = render_latent
# unianimate slices
partial_unianim_data = None
if unianim_data is not None:
partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx]
partial_unianim_data = {
"dwpose": partial_dwpose,
"random_ref": unianim_data["random_ref"],
"strength": unianimate_poses["strength"],
"start_percent": unianimate_poses["start_percent"],
"end_percent": unianimate_poses["end_percent"]
}
# fantasy portrait slices
partial_fantasy_portrait_input = None
if fantasy_portrait_input is not None:
adapter_proj = fantasy_portrait_input["adapter_proj"]
if latent_end_idx > adapter_proj.shape[1]:
pad_len = latent_end_idx - adapter_proj.shape[1]
last_frame = adapter_proj[:, -1:, :, :].repeat(1, pad_len, 1, 1)
padded_proj = torch.cat([adapter_proj, last_frame], dim=1)
else:
padded_proj = adapter_proj
partial_fantasy_portrait_input = fantasy_portrait_input.copy()
partial_fantasy_portrait_input["adapter_proj"] = padded_proj[:, latent_start_idx:latent_end_idx]
mm.soft_empty_cache()
gc.collect()
# sampling loop
sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True)
for i in range(len(timesteps)-1):
timestep = timesteps[i]
latent_model_input = latent.to(device)
if mode == "infinitetalk":
if humo_image_cond is None or not is_first_clip:
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
noise_pred, _, self.cache_state = predict_func(
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
timestep, i, y, clip_embeds, control_latents, None, 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,
humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg,
uni3c_data = uni3c_data)
if callback is not None:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * timestep.to(device) / 1000).detach().permute(1,0,2,3)
callback(step_iteration_count, callback_latent, None, estimated_iterations*(len(timesteps)-1))
del callback_latent
sampling_pbar.update(1)
step_iteration_count += 1
# update latent
if use_tsr:
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
if scheduler == "multitalk":
noise_pred = -noise_pred
dt = (timesteps[i] - timesteps[i + 1]) / 1000
latent = latent + noise_pred * dt[:, None, None, None]
else:
latent = sample_scheduler.step(noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0).to(noise_pred.device), **scheduler_step_args)[0].squeeze(0)
del noise_pred, latent_model_input, timestep
# differential diffusion inpaint
if masks is not None:
if i < len(timesteps) - 1:
image_latent = add_noise(original_image.to(device), noise.to(device), timesteps[i+1])
mask = masks[i].to(latent)
latent = image_latent * mask + latent * (1-mask)
# injecting motion frames
if not is_first_clip and mode == "multitalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
latent[:, :add_latent.shape[1]] = add_latent
else:
if humo_image_cond is None or not is_first_clip:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
del noise, latent_motion_frames
if offload:
offload_transformer(transformer, remove_lora=False)
offloaded = True
if humo_image_cond is not None and humo_reference_count > 0:
latent = latent[:,:-humo_reference_count]
vae.to(device)
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
vae.to(offload_device)
sampling_pbar.close()
# optional color correction (less relevant for InfiniteTalk)
if colormatch != "disabled":
videos = videos.permute(1, 2, 3, 0).float().numpy()
from color_matcher import ColorMatcher
cm = ColorMatcher()
cm_result_list = []
for img in videos:
if mode == "multitalk":
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
else:
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype))
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
# optionally save generated samples to disk
if output_path:
video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8')
num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num
log.info(f"Saving {num_frames_to_save} generated frames to {output_path}")
start_idx = 0 if is_first_clip else cur_motion_frames_num
for i in range(start_idx, video_np.shape[0]):
im = Image.fromarray(video_np[i])
im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png"))
img_counter += 1
else:
gen_video_list.append(videos if is_first_clip else videos[:, cur_motion_frames_num:])
current_condframe_index += 1
iteration_count += 1
# decide whether is done
if arrive_last_frame:
break
# update next condition frames
is_first_clip = False
cur_motion_frames_num = motion_frame
cond_ = videos[:, -cur_motion_frames_num:].unsqueeze(0)
if mode == "infinitetalk":
cond_frame = cond_
else:
cond_image = cond_
del videos, latent
# Repeat audio emb
if multitalk_embeds is not None:
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count)
audio_end_idx = audio_start_idx + clip_length
if audio_end_idx >= len(audio_embedding[0]):
arrive_last_frame = True
miss_lengths = []
source_frames = []
for human_inx in range(human_num):
source_frame = len(audio_embedding[human_inx])
source_frames.append(source_frame)
if audio_end_idx >= len(audio_embedding[human_inx]):
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
log.warning(f"Padding length: {miss_length}")
if encoded_silence is not None:
add_audio_emb = encoded_silence[-1*miss_length:]
else:
add_audio_emb = torch.flip(audio_embedding[human_inx][-1*miss_length:], dims=[0])
audio_embedding[human_inx] = torch.cat([audio_embedding[human_inx], add_audio_emb.to(device, dtype)], dim=0)
miss_lengths.append(miss_length)
else:
miss_lengths.append(0)
if mode == "infinitetalk" and current_condframe_index >= original_images.shape[2]:
last_frame = original_images[:, :, -1:, :, :]
miss_length = 1
original_images = torch.cat([original_images, last_frame.repeat(1, 1, miss_length, 1, 1)], dim=2)
if not output_path:
gen_video_samples = torch.cat(gen_video_list, dim=1)
else:
gen_video_samples = torch.zeros(3, 1, 64, 64) # dummy output
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.permute(1, 2, 3, 0), "output_path": output_path},
-28
View File
@@ -1848,32 +1848,6 @@ class WanVideoContextOptions:
return (context_options,)
class WanVideoFlowEdit:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"source_embeds": ("WANVIDEOTEXTEMBEDS", ),
"skip_steps": ("INT", {"default": 4, "min": 0}),
"drift_steps": ("INT", {"default": 0, "min": 0}),
"drift_flow_shift": ("FLOAT", {"default": 3.0, "min": 1.0, "max": 30.0, "step": 0.01}),
"source_cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
"drift_cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
},
"optional": {
"source_image_embeds": ("WANVIDIMAGE_EMBEDS", ),
}
}
RETURN_TYPES = ("FLOWEDITARGS", )
RETURN_NAMES = ("flowedit_args",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Flowedit options for WanVideo"
def process(self, **kwargs):
return (kwargs,)
class WanVideoLoopArgs:
@classmethod
def INPUT_TYPES(s):
@@ -2249,7 +2223,6 @@ NODE_CLASS_MAPPINGS = {
"WanVideoEnhanceAVideo": WanVideoEnhanceAVideo,
"WanVideoContextOptions": WanVideoContextOptions,
"WanVideoTextEmbedBridge": WanVideoTextEmbedBridge,
"WanVideoFlowEdit": WanVideoFlowEdit,
"WanVideoControlEmbeds": WanVideoControlEmbeds,
"WanVideoSLG": WanVideoSLG,
"WanVideoLoopArgs": WanVideoLoopArgs,
@@ -2292,7 +2265,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video",
"WanVideoContextOptions": "WanVideo Context Options",
"WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge",
"WanVideoFlowEdit": "WanVideo FlowEdit",
"WanVideoControlEmbeds": "WanVideo Control Embeds",
"WanVideoSLG": "WanVideo SLG",
"WanVideoLoopArgs": "WanVideo Loop Args",
+104 -787
View File
@@ -3,16 +3,14 @@ import torch
import numpy as np
from tqdm import tqdm
import inspect
from PIL import Image
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from .wanvideo.schedulers.fm_solvers import get_sampling_sigmas, retrieve_timesteps
from .wanvideo.modules.model import rope_params
from .custom_linear import remove_lora_from_module, set_lora_params, _replace_linear
from .wanvideo.schedulers import get_scheduler, scheduler_list
from .gguf.gguf import set_lora_params_gguf
from .multitalk.multitalk import timestep_transform, add_noise
from .multitalk.multitalk import add_noise
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)
compile_model, dict_to_device, tangential_projection, get_raag_guidance, temporal_score_rescaling, offload_transformer, init_blockswap)
from .multitalk.multitalk_loop import multitalk_loop
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
@@ -20,7 +18,7 @@ from .WanMove.trajectory import replace_feature
from contextlib import nullcontext
from comfy import model_management as mm
from comfy.utils import ProgressBar, load_torch_file
from comfy.utils import ProgressBar
from comfy.cli_args import args, LatentPreviewMethod
script_directory = os.path.dirname(os.path.abspath(__file__))
@@ -33,82 +31,6 @@ rope_functions = ["default", "comfy", "comfy_chunked"]
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
try:
from .gguf.gguf import GGUFParameter
except:
pass
class MetaParameter(torch.nn.Parameter):
def __new__(cls, dtype, quant_type=None):
data = torch.empty(0, dtype=dtype)
self = torch.nn.Parameter(data, requires_grad=False)
self.quant_type = quant_type
return self
def offload_transformer(transformer, remove_lora=True):
transformer.teacache_state.clear_all()
transformer.magcache_state.clear_all()
transformer.easycache_state.clear_all()
if transformer.patched_linear:
for name, param in transformer.named_parameters():
if "loras" in name or "controlnet" in name:
continue
module = transformer
subnames = name.split('.')
for subname in subnames[:-1]:
module = getattr(module, subname)
attr_name = subnames[-1]
if param.data.is_floating_point():
meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False)
setattr(module, attr_name, meta_param)
elif isinstance(param.data, GGUFParameter):
quant_type = getattr(param, 'quant_type', None)
setattr(module, attr_name, MetaParameter(param.data.dtype, quant_type))
else:
pass
if remove_lora:
remove_lora_from_module(transformer)
else:
transformer.to(offload_device)
for block in transformer.blocks:
block.kv_cache = None
if transformer.audio_model is not None and hasattr(block, 'audio_block'):
block.audio_block = None
mm.soft_empty_cache()
gc.collect()
def init_blockswap(transformer, block_swap_args, model):
if not transformer.patched_linear:
if block_swap_args is not None:
for name, param in transformer.named_parameters():
if "block" not in name or "control_adapter" in name or "face" 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()
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)
class WanVideoSampler:
@classmethod
@@ -132,7 +54,7 @@ class WanVideoSampler:
"feta_args": ("FETAARGS", ),
"context_options": ("WANVIDCONTEXT", ),
"cache_args": ("CACHEARGS", ),
"flowedit_args": ("FLOWEDITARGS", ),
"flowedit_args": ("FLOWEDITARGS", {"tooltip": "FlowEdit support has been deprecated"}),
"batched_cfg": ("BOOLEAN", {"default": False, "tooltip": "Batch cond and uncond for faster sampling, possibly faster on some hardware, uses more memory"}),
"slg_args": ("SLGARGS", ),
"rope_function": (rope_functions, {"default": "comfy", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile. Chunked version has reduced peak VRAM usage when not using torch.compile"}),
@@ -159,7 +81,8 @@ class WanVideoSampler:
force_offload=True, samples=None, feta_args=None, denoise_strength=1.0, context_options=None,
cache_args=None, teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None,
experimental_args=None, sigmas=None, unianimate_poses=None, fantasytalking_embeds=None, uni3c_embeds=None, multitalk_embeds=None, freeinit_args=None, start_step=0, end_step=-1, add_noise_to_samples=False):
if flowedit_args is not None:
raise Exception("FlowEdit support has been deprecated and removed due to lack of use and code maintainability")
patcher = model
model = model.model
transformer = model.diffusion_model
@@ -253,7 +176,7 @@ class WanVideoSampler:
timesteps = scheduler["timesteps"]
start_step = scheduler.get("start_step", start_step)
elif scheduler != "multitalk":
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, log_timesteps=True)
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, denoise_strength, sigmas=sigmas, log_timesteps=True)
log.info(f"sigmas: {sample_scheduler.sigmas}")
else:
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
@@ -994,43 +917,6 @@ class WanVideoSampler:
if transformer.attention_mode == "radial_sage_attention":
setup_radial_attention(transformer, transformer_options, latent, seq_len, latent_video_length, context_options=context_options)
# FlowEdit setup
if flowedit_args is not None:
source_embeds = flowedit_args["source_embeds"]
source_embeds = dict_to_device(source_embeds, device)
source_image_embeds = flowedit_args.get("source_image_embeds", image_embeds)
source_image_cond = source_image_embeds.get("image_embeds", None)
source_clip_fea = source_image_embeds.get("clip_fea", clip_fea)
if source_image_cond is not None:
source_image_cond = source_image_cond.to(dtype)
skip_steps = flowedit_args["skip_steps"]
drift_steps = flowedit_args["drift_steps"]
source_cfg = flowedit_args["source_cfg"]
if not isinstance(source_cfg, list):
source_cfg = [source_cfg] * (steps +1)
drift_cfg = flowedit_args["drift_cfg"]
if not isinstance(drift_cfg, list):
drift_cfg = [drift_cfg] * (steps +1)
x_init = samples["samples"].clone().squeeze(0).to(device)
x_tgt = samples["samples"].squeeze(0).to(device)
sample_scheduler = FlowMatchEulerDiscreteScheduler(
num_train_timesteps=1000,
shift=flowedit_args["drift_flow_shift"],
use_dynamic_shifting=False)
sampling_sigmas = get_sampling_sigmas(steps, flowedit_args["drift_flow_shift"])
drift_timesteps, _ = retrieve_timesteps(
sample_scheduler,
device=device,
sigmas=sampling_sigmas)
if drift_steps > 0:
drift_timesteps = torch.cat([drift_timesteps, torch.tensor([0]).to(drift_timesteps.device)]).to(drift_timesteps.device)
timesteps[-drift_steps:] = drift_timesteps[-drift_steps:]
# Experimental args
use_cfg_zero_star = use_tangential = use_fresca = bidirectional_sampling = use_tsr = False
raag_alpha = 0.0
@@ -1834,7 +1720,7 @@ class WanVideoSampler:
# FreeInit noise reinitialization (after first iteration)
if freeinit_args is not None and iter_idx > 0:
# restart scheduler for each iteration
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, denoise_strength, sigmas=sigmas)
# Re-apply start_step and end_step logic to timesteps and sigmas
if end_step != -1:
@@ -1884,9 +1770,6 @@ class WanVideoSampler:
pbar = ProgressBar(len(timesteps) - ttm_start_step)
#region main loop start
for idx, t in enumerate(tqdm(timesteps[ttm_start_step:], disable=multitalk_sampling or wananimate_loop)):
if flowedit_args is not None:
if idx < skip_steps:
continue
if bidirectional_sampling:
latent_flipped = torch.flip(latent, dims=[1])
@@ -1941,129 +1824,6 @@ class WanVideoSampler:
enhance_enabled = False
if feta_args is not None and feta_start_percent <= current_step_percentage <= feta_end_percent:
enhance_enabled = True
#flow-edit
if flowedit_args is not None:
sigma = t / 1000.0
sigma_prev = (timesteps[idx + 1] if idx < len(timesteps) - 1 else timesteps[-1]) / 1000.0
noise = torch.randn(x_init.shape, generator=seed_g, device=torch.device("cpu"))
if idx < len(timesteps) - drift_steps:
cfg = drift_cfg
zt_src = (1-sigma) * x_init + sigma * noise.to(t)
zt_tgt = x_tgt + zt_src - x_init
#source
if idx < len(timesteps) - drift_steps:
if context_options is not None:
counter = torch.zeros_like(zt_src, device=intermediate_device)
vt_src = torch.zeros_like(zt_src, device=intermediate_device)
context_queue = list(context(idx, steps, latent_video_length, context_frames, context_stride, context_overlap))
for c in context_queue:
window_id = self.window_tracker.get_window_id(c)
if cache_args is not None:
current_teacache = self.window_tracker.get_teacache(window_id, self.cache_state)
else:
current_teacache = None
prompt_index = min(int(max(c) / section_size), num_prompts - 1)
if context_options["verbose"]:
log.info(f"Prompt index: {prompt_index}")
if len(source_embeds["prompt_embeds"]) > 1:
positive = source_embeds["prompt_embeds"][prompt_index]
else:
positive = source_embeds["prompt_embeds"]
partial_img_emb = None
if source_image_cond is not None:
partial_img_emb = source_image_cond[:, c, :, :]
partial_img_emb[:, 0, :, :] = source_image_cond[:, 0, :, :].to(intermediate_device)
partial_zt_src = zt_src[:, c, :, :]
vt_src_context, _, new_teacache = predict_with_cfg(
partial_zt_src, cfg[idx],
positive, source_embeds["negative_prompt_embeds"],
timestep, idx, partial_img_emb, control_latents,
source_clip_fea, current_teacache)
if cache_args is not None:
self.window_tracker.cache_states[window_id] = new_teacache
window_mask = create_window_mask(vt_src_context, c, latent_video_length, context_overlap)
vt_src[:, c, :, :] += vt_src_context * window_mask
counter[:, c, :, :] += window_mask
vt_src /= counter
else:
vt_src, _, self.cache_state_source = predict_with_cfg(
zt_src, cfg[idx],
source_embeds["prompt_embeds"],
source_embeds["negative_prompt_embeds"],
timestep, idx, source_image_cond,
source_clip_fea, control_latents,
cache_state=self.cache_state_source)
else:
if idx == len(timesteps) - drift_steps:
x_tgt = zt_tgt
zt_tgt = x_tgt
vt_src = 0
#target
if context_options is not None:
counter = torch.zeros_like(zt_tgt, device=intermediate_device)
vt_tgt = torch.zeros_like(zt_tgt, device=intermediate_device)
context_queue = list(context(idx, steps, latent_video_length, context_frames, context_stride, context_overlap))
for c in context_queue:
window_id = self.window_tracker.get_window_id(c)
if cache_args is not None:
current_teacache = self.window_tracker.get_teacache(window_id, self.cache_state)
else:
current_teacache = None
prompt_index = min(int(max(c) / section_size), num_prompts - 1)
if context_options["verbose"]:
log.info(f"Prompt index: {prompt_index}")
if len(text_embeds["prompt_embeds"]) > 1:
positive = text_embeds["prompt_embeds"][prompt_index]
else:
positive = text_embeds["prompt_embeds"]
partial_img_emb = None
partial_control_latents = None
if image_cond is not None:
partial_img_emb = image_cond[:, c, :, :]
partial_img_emb[:, 0, :, :] = image_cond[:, 0, :, :].to(intermediate_device)
if control_latents is not None:
partial_control_latents = control_latents[:, c, :, :]
partial_zt_tgt = zt_tgt[:, c, :, :]
vt_tgt_context, _, new_teacache = predict_with_cfg(
partial_zt_tgt, cfg[idx],
positive, text_embeds["negative_prompt_embeds"],
timestep, idx, partial_img_emb, partial_control_latents,
clip_fea, current_teacache)
if cache_args is not None:
self.window_tracker.cache_states[window_id] = new_teacache
window_mask = create_window_mask(vt_tgt_context, c, latent_video_length, context_overlap)
vt_tgt[:, c, :, :] += vt_tgt_context * window_mask
counter[:, c, :, :] += window_mask
vt_tgt /= counter
else:
vt_tgt, _,self.cache_state = predict_with_cfg(
zt_tgt, cfg[idx],
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, idx, image_cond, clip_fea, control_latents,
cache_state=self.cache_state)
v_delta = vt_tgt - vt_src
x_tgt = x_tgt.to(torch.float32)
v_delta = v_delta.to(torch.float32)
x_tgt = x_tgt + (sigma_prev - sigma) * v_delta
x0 = x_tgt
#region context windowing
elif context_options is not None:
counter = torch.zeros_like(latent_model_input, device=intermediate_device)
@@ -2258,443 +2018,7 @@ class WanVideoSampler:
noise_pred /= counter
#region multitalk
elif multitalk_sampling:
mode = image_embeds.get("multitalk_mode", "multitalk")
if mode == "auto":
mode = transformer.multitalk_model_type.lower()
log.info(f"Multitalk mode: {mode}")
cond_frame = None
offload = image_embeds.get("force_offload", False)
offloaded = False
tiled_vae = image_embeds.get("tiled_vae", False)
frame_num = clip_length = image_embeds.get("frame_window_size", 81)
clip_embeds = image_embeds.get("clip_context", None)
if clip_embeds is not None:
clip_embeds = clip_embeds.to(dtype)
colormatch = image_embeds.get("colormatch", "disabled")
motion_frame = image_embeds.get("motion_frame", 25)
target_w = image_embeds.get("target_w", None)
target_h = image_embeds.get("target_h", None)
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
if original_images is None:
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
output_path = image_embeds.get("output_path", "")
img_counter = 0
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
face_scale = 0.1
x_min, x_max = int(target_h * face_scale), int(target_h * (1 - face_scale))
lefty_min, lefty_max = int((target_w//2) * face_scale), int((target_w//2) * (1 - face_scale))
righty_min, righty_max = int((target_w//2) * face_scale + (target_w//2)), int((target_w//2) * (1 - face_scale) + (target_w//2))
human_mask1, human_mask2 = (torch.zeros([target_h, target_w]) for _ in range(2))
human_mask1[x_min:x_max, lefty_min:lefty_max] = 1
human_mask2[x_min:x_max, righty_min:righty_max] = 1
background_mask = torch.where((human_mask1 + human_mask2) > 0, torch.tensor(0), torch.tensor(1))
human_masks = [human_mask1, human_mask2, background_mask]
ref_target_masks = torch.stack(human_masks, dim=0)
multitalk_embeds['ref_target_masks'] = ref_target_masks
gen_video_list = []
is_first_clip = True
arrive_last_frame = False
cur_motion_frames_num = 1
audio_start_idx = iteration_count = step_iteration_count = 0
audio_end_idx = (audio_start_idx + clip_length) * audio_stride
indices = (torch.arange(4 + 1) - 2) * 1
current_condframe_index = 0
audio_embedding = multitalk_audio_embeds
human_num = len(audio_embedding)
audio_embs = None
cond_frame = None
uni3c_data = uni3c_data_input = None
if uni3c_embeds is not None:
transformer.controlnet = uni3c_embeds["controlnet"]
uni3c_data = uni3c_embeds.copy()
encoded_silence = None
try:
silence_path = os.path.join(script_directory, "multitalk", "encoded_silence.safetensors")
encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype)
except:
log.warning("No encoded silence file found, padding with end of audio embedding instead.")
total_frames = len(audio_embedding[0])
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
callback = prepare_callback(patcher, estimated_iterations)
if frame_num >= total_frames:
arrive_last_frame = True
estimated_iterations = 1
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
while True: # start video generation iteratively
self.cache_state = [None, None]
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
if mode == "infinitetalk":
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None
if multitalk_embeds is not None:
audio_embs = []
# split audio with window size
for human_idx in range(human_num):
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0]-1)
audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device)
audio_embs.append(audio_emb)
audio_embs = torch.concat(audio_embs, dim=0).to(dtype)
h, w = (cond_image.shape[-2], cond_image.shape[-1]) if cond_image is not None else (target_h, target_w)
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
seq_len = ((frame_num - 1) // VAE_STRIDE[0] + 1) * lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
latent_frame_num = (frame_num - 1) // 4 + 1
noise = torch.randn(
16, latent_frame_num,
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
# Calculate the correct latent slice based on current iteration
if is_first_clip:
latent_start_idx = 0
latent_end_idx = noise.shape[1]
else:
new_frames_per_iteration = frame_num - motion_frame
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
latent_start_idx = iteration_count * new_latent_frames_per_iteration
latent_end_idx = latent_start_idx + noise.shape[1]
if samples is not None:
noise_mask = samples.get("noise_mask", None)
input_samples = samples["samples"]
if input_samples is not None:
input_samples = input_samples.squeeze(0).to(noise)
# Check if we have enough frames in input_samples
if latent_end_idx > input_samples.shape[1]:
# We need more frames than available - pad the input_samples at the end
pad_length = latent_end_idx - input_samples.shape[1]
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
input_samples = torch.cat([input_samples, last_frame], dim=1)
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
if noise_mask is not None:
original_image = input_samples.to(device)
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
if add_noise_to_samples:
latent_timestep = timesteps[0]
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
else:
noise = input_samples
# diff diff prep
if noise_mask is not None:
if len(noise_mask.shape) == 4:
noise_mask = noise_mask.squeeze(1)
if audio_end_idx > noise_mask.shape[0]:
noise_mask = noise_mask.repeat(audio_end_idx // noise_mask.shape[0], 1, 1)
noise_mask = noise_mask[audio_start_idx:audio_end_idx]
noise_mask = torch.nn.functional.interpolate(
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
mode='trilinear',
align_corners=False
).repeat(1, noise.shape[0], 1, 1, 1)
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device)
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
# zero padding and vae encode for img cond
if cond_image is not None or cond_frame is not None:
cond_ = cond_image if (is_first_clip or humo_image_cond is None) else cond_frame
cond_frame_num = cond_.shape[2]
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num, target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.concat([cond_.to(device, vae.dtype), video_frames], dim=2)
# encode
vae.to(device)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
if mode == "multitalk":
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
else:
cond_ = cond_image if is_first_clip else cond_frame
latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
vae.to(offload_device)
#motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
msk[:, :1] = 1
y = torch.cat([msk, y]) # 4+C T H W
mm.soft_empty_cache()
else:
y = None
latent_motion_frames = noise[:, :1]
partial_humo_cond_input = partial_humo_cond_neg_input = partial_humo_audio = partial_humo_audio_neg = None
if humo_image_cond is not None:
partial_humo_cond_input = humo_image_cond[:, :latent_frame_num]
partial_humo_cond_neg_input = humo_image_cond_neg[:, :latent_frame_num]
if y is not None:
partial_humo_cond_input[:, :1] = y[:, :1]
if humo_reference_count > 0:
partial_humo_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:]
partial_humo_cond_neg_input[:, -humo_reference_count:] = humo_image_cond_neg[:, -humo_reference_count:]
if humo_audio is not None:
if is_first_clip:
audio_embs = None
partial_humo_audio, _ = get_audio_emb_window(humo_audio, frame_num, frame0_idx=audio_start_idx)
#zero_audio_pad = torch.zeros(humo_reference_count, *partial_humo_audio.shape[1:], device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
partial_humo_audio[-humo_reference_count:] = 0
partial_humo_audio_neg = torch.zeros_like(partial_humo_audio, device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
if scheduler == "multitalk":
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
timesteps.append(0.)
timesteps = [torch.tensor([t], device=device) for t in timesteps]
timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps]
else:
if isinstance(scheduler, dict):
sample_scheduler = copy.deepcopy(scheduler["sample_scheduler"])
timesteps = scheduler["timesteps"]
else:
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)]
# sample videos
latent = noise
# injecting motion frames
if not is_first_clip and mode == "multitalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
latent[:, :add_latent.shape[1]] = add_latent
if offloaded:
# 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)
elif 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)
#blockswap init
init_blockswap(transformer, block_swap_args, model)
# Use the appropriate prompt for this section
if len(text_embeds["prompt_embeds"]) > 1:
prompt_index = min(iteration_count, len(text_embeds["prompt_embeds"]) - 1)
positive = [text_embeds["prompt_embeds"][prompt_index]]
log.info(f"Using prompt index: {prompt_index}")
else:
positive = text_embeds["prompt_embeds"]
# uni3c slices
if uni3c_embeds is not None:
vae.to(device)
# Pad original_images if needed
num_frames = original_images.shape[2]
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)
padded_images = torch.cat([original_images, last_frame], dim=2)
else:
padded_images = original_images
render_latent = vae.encode(
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
device=device, tiled=tiled_vae
).to(dtype)
vae.to(offload_device)
uni3c_data['render_latent'] = render_latent
# unianimate slices
partial_unianim_data = None
if unianim_data is not None:
partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx]
partial_unianim_data = {
"dwpose": partial_dwpose,
"random_ref": unianim_data["random_ref"],
"strength": unianimate_poses["strength"],
"start_percent": unianimate_poses["start_percent"],
"end_percent": unianimate_poses["end_percent"]
}
# fantasy portrait slices
partial_fantasy_portrait_input = None
if fantasy_portrait_input is not None:
adapter_proj = fantasy_portrait_input["adapter_proj"]
if latent_end_idx > adapter_proj.shape[1]:
pad_len = latent_end_idx - adapter_proj.shape[1]
last_frame = adapter_proj[:, -1:, :, :].repeat(1, pad_len, 1, 1)
padded_proj = torch.cat([adapter_proj, last_frame], dim=1)
else:
padded_proj = adapter_proj
partial_fantasy_portrait_input = fantasy_portrait_input.copy()
partial_fantasy_portrait_input["adapter_proj"] = padded_proj[:, latent_start_idx:latent_end_idx]
mm.soft_empty_cache()
gc.collect()
# sampling loop
sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True)
for i in range(len(timesteps)-1):
timestep = timesteps[i]
latent_model_input = latent.to(device)
if mode == "infinitetalk":
if humo_image_cond is None or not is_first_clip:
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
noise_pred, _, self.cache_state = predict_with_cfg(
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
timestep, i, y, clip_embeds, control_latents, None, 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,
humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg,
uni3c_data = uni3c_data)
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, estimated_iterations*(len(timesteps)-1))
del callback_latent
sampling_pbar.update(1)
step_iteration_count += 1
# update latent
if use_tsr:
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
if scheduler == "multitalk":
noise_pred = -noise_pred
dt = (timesteps[i] - timesteps[i + 1]) / 1000
latent = latent + noise_pred * dt[:, None, None, None]
else:
latent = sample_scheduler.step(noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0).to(noise_pred.device), **scheduler_step_args)[0].squeeze(0)
del noise_pred, latent_model_input, timestep
# differential diffusion inpaint
if masks is not None:
if i < len(timesteps) - 1:
image_latent = add_noise(original_image.to(device), noise.to(device), timesteps[i+1])
mask = masks[i].to(latent)
latent = image_latent * mask + latent * (1-mask)
# injecting motion frames
if not is_first_clip and mode == "multitalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
latent[:, :add_latent.shape[1]] = add_latent
else:
if humo_image_cond is None or not is_first_clip:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
del noise, latent_motion_frames
if offload:
offload_transformer(transformer, remove_lora=False)
offloaded = True
if humo_image_cond is not None and humo_reference_count > 0:
latent = latent[:,:-humo_reference_count]
vae.to(device)
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
vae.to(offload_device)
sampling_pbar.close()
# optional color correction (less relevant for InfiniteTalk)
if colormatch != "disabled":
videos = videos.permute(1, 2, 3, 0).float().numpy()
from color_matcher import ColorMatcher
cm = ColorMatcher()
cm_result_list = []
for img in videos:
if mode == "multitalk":
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
else:
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype))
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
# optionally save generated samples to disk
if output_path:
video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8')
num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num
log.info(f"Saving {num_frames_to_save} generated frames to {output_path}")
start_idx = 0 if is_first_clip else cur_motion_frames_num
for i in range(start_idx, video_np.shape[0]):
im = Image.fromarray(video_np[i])
im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png"))
img_counter += 1
else:
gen_video_list.append(videos if is_first_clip else videos[:, cur_motion_frames_num:])
current_condframe_index += 1
iteration_count += 1
# decide whether is done
if arrive_last_frame:
break
# update next condition frames
is_first_clip = False
cur_motion_frames_num = motion_frame
cond_ = videos[:, -cur_motion_frames_num:].unsqueeze(0)
if mode == "infinitetalk":
cond_frame = cond_
else:
cond_image = cond_
del videos, latent
# Repeat audio emb
if multitalk_embeds is not None:
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count)
audio_end_idx = audio_start_idx + clip_length
if audio_end_idx >= len(audio_embedding[0]):
arrive_last_frame = True
miss_lengths = []
source_frames = []
for human_inx in range(human_num):
source_frame = len(audio_embedding[human_inx])
source_frames.append(source_frame)
if audio_end_idx >= len(audio_embedding[human_inx]):
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
log.warning(f"Padding length: {miss_length}")
if encoded_silence is not None:
add_audio_emb = encoded_silence[-1*miss_length:]
else:
add_audio_emb = torch.flip(audio_embedding[human_inx][-1*miss_length:], dims=[0])
audio_embedding[human_inx] = torch.cat([audio_embedding[human_inx], add_audio_emb.to(device, dtype)], dim=0)
miss_lengths.append(miss_length)
else:
miss_lengths.append(0)
if mode == "infinitetalk" and current_condframe_index >= original_images.shape[2]:
last_frame = original_images[:, :, -1:, :, :]
miss_length = 1
original_images = torch.cat([original_images, last_frame.repeat(1, 1, miss_length, 1, 1)], dim=2)
if not output_path:
gen_video_samples = torch.cat(gen_video_list, dim=1)
else:
gen_video_samples = torch.zeros(3, 1, 64, 64) # dummy output
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.permute(1, 2, 3, 0), "output_path": output_path},
return multitalk_loop(**locals())
# region framepack loop
elif framepack:
framepack_out = []
@@ -2779,7 +2103,7 @@ class WanVideoSampler:
sample_scheduler = copy.deepcopy(scheduler["sample_scheduler"])
timesteps = scheduler["timesteps"]
else:
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, 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)):
@@ -2959,12 +2283,12 @@ class WanVideoSampler:
if input_samples is not None:
input_samples = input_samples.squeeze(0).to(noise)
# Check if we have enough frames in input_samples
if latent_end_idx > input_samples.shape[1]:
# We need more frames than available - pad the input_samples at the end
pad_length = latent_end_idx - input_samples.shape[1]
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
input_samples = torch.cat([input_samples, last_frame], dim=1)
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
# if latent_end_idx > input_samples.shape[1]:
# # We need more frames than available - pad the input_samples at the end
# pad_length = latent_end_idx - input_samples.shape[1]
# last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
# input_samples = torch.cat([input_samples, last_frame], dim=1)
# input_samples = input_samples[:, latent_start_idx:latent_end_idx]
if noise_mask is not None:
original_image = input_samples.to(device)
@@ -3000,7 +2324,7 @@ class WanVideoSampler:
sample_scheduler = copy.deepcopy(scheduler["sample_scheduler"])
timesteps = scheduler["timesteps"]
else:
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, denoise_strength, sigmas=sigmas)
# sample videos
latent = noise
@@ -3096,17 +2420,17 @@ class WanVideoSampler:
current_ref_images = videos[:, -refert_num:].clone().detach()
# optionally save generated samples to disk
if output_path:
video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8')
num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num
log.info(f"Saving {num_frames_to_save} generated frames to {output_path}")
start_idx = 0 if is_first_clip else cur_motion_frames_num
for i in range(start_idx, video_np.shape[0]):
im = Image.fromarray(video_np[i])
im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png"))
img_counter += 1
else:
gen_video_list.append(videos)
# if output_path:
# video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8')
# num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num
# log.info(f"Saving {num_frames_to_save} generated frames to {output_path}")
# start_idx = 0 if is_first_clip else cur_motion_frames_num
# for i in range(start_idx, video_np.shape[0]):
# im = Image.fromarray(video_np[i])
# im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png"))
# img_counter += 1
# else:
gen_video_list.append(videos)
del videos
@@ -3155,94 +2479,87 @@ class WanVideoSampler:
noise_pred = torch.cat([noise_pred[:, latent_video_length - shift_idx:]] + [noise_pred[:, :latent_video_length - shift_idx]], dim=1)
shift_idx = (shift_idx + latent_skip) % latent_video_length
latent = latent.to(intermediate_device)
if flowedit_args is None:
latent = latent.to(intermediate_device)
if self.noise_front_pad_num > 0:
noise_pred = noise_pred[:, self.noise_front_pad_num:]
if self.noise_front_pad_num > 0:
noise_pred = noise_pred[:, self.noise_front_pad_num:]
if use_tsr:
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
if use_tsr:
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
if transformer.is_longcat:
noise_pred = -noise_pred
if transformer.is_longcat:
noise_pred = -noise_pred
if len(timestep.shape) != 1 and clean_latent_indices and not is_pusa: #5b and longcat, skip clean latents for scheduler step
step_process_indices = [i for i in range(latent.shape[1]) if i not in clean_latent_indices]
latent[:, step_process_indices] = sample_scheduler.step(noise_pred[:, step_process_indices].unsqueeze(0), orig_timestep,
latent[:, step_process_indices].unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
else:
if latents_to_not_step > 0:
raw_latent = latent[:, :latents_to_not_step]
noise_pred_in = noise_pred[:, latents_to_not_step:]
latent = latent[:, latents_to_not_step:]
elif recammaster is not None or mocha_embeds is not None:
noise_pred_in = noise_pred[:, :orig_noise_len]
latent = latent[:, :orig_noise_len]
else:
noise_pred_in = noise_pred
latent = sample_scheduler.step(noise_pred_in.unsqueeze(0), timestep, latent.unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
if noise_pred_flipped is not None:
latent_backwards = sample_scheduler_flipped.step(noise_pred_flipped.unsqueeze(0), timestep, latent_flipped.unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
latent_backwards = torch.flip(latent_backwards, dims=[1])
latent = latent * 0.5 + latent_backwards * 0.5
if latents_to_not_step > 0:
latent = torch.cat([raw_latent, latent], dim=1)
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)
#InfiniteTalk first frame handling
if (extra_latents is not None
and not multitalk_sampling
and transformer.multitalk_model_type=="InfiniteTalk"):
for entry in extra_latents:
add_index = entry["index"]
num_extra_frames = entry["samples"].shape[2]
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
# differential diffusion inpaint
if masks is not None:
if idx < len(timesteps) - 1:
noise_timestep = timesteps[idx+1]
image_latent = sample_scheduler.scale_noise(
original_image.to(device), torch.tensor([noise_timestep]), noise.to(device)
)
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):
noisy_latents = add_noise(ttm_reference_latents, noise, sample_scheduler.all_timesteps[idx + ttm_start_step + 1].to(noise.device)).to(latent)
latent = latent * (1 - motion_mask) + noisy_latents * motion_mask
else:
latent = latent * (1 - motion_mask) + ttm_reference_latents.to(latent) * motion_mask
if freeinit_args is not None:
current_latent = latent.clone()
if callback is not None:
if recammaster is not None or mocha_embeds is not None:
callback_latent = (latent_model_input[:, :orig_noise_len].to(device) - noise_pred[:, :orig_noise_len].to(device) * t.to(device) / 1000).detach()
#elif phantom_latents is not None:
# callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach()
elif humo_reference_count > 0:
callback_latent = (latent_model_input[:,:-humo_reference_count].to(device) - noise_pred[:,:-humo_reference_count].to(device) * t.to(device) / 1000).detach()
elif "rcm" in sample_scheduler.__class__.__name__.lower():
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device)).detach()
else:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach()
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
else:
pbar.update(1)
if len(timestep.shape) != 1 and clean_latent_indices and not is_pusa: #5b and longcat, skip clean latents for scheduler step
step_process_indices = [i for i in range(latent.shape[1]) if i not in clean_latent_indices]
latent[:, step_process_indices] = sample_scheduler.step(noise_pred[:, step_process_indices].unsqueeze(0), orig_timestep,
latent[:, step_process_indices].unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
else:
if callback is not None:
callback_latent = (zt_tgt.to(device) - vt_tgt.to(device) * t.to(device) / 1000).detach()
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
if latents_to_not_step > 0:
raw_latent = latent[:, :latents_to_not_step]
noise_pred_in = noise_pred[:, latents_to_not_step:]
latent = latent[:, latents_to_not_step:]
elif recammaster is not None or mocha_embeds is not None:
noise_pred_in = noise_pred[:, :orig_noise_len]
latent = latent[:, :orig_noise_len]
else:
pbar.update(1)
noise_pred_in = noise_pred
latent = sample_scheduler.step(noise_pred_in.unsqueeze(0), timestep, latent.unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
if noise_pred_flipped is not None:
latent_backwards = sample_scheduler_flipped.step(noise_pred_flipped.unsqueeze(0), timestep, latent_flipped.unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
latent_backwards = torch.flip(latent_backwards, dims=[1])
latent = latent * 0.5 + latent_backwards * 0.5
if latents_to_not_step > 0:
latent = torch.cat([raw_latent, latent], dim=1)
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)
#InfiniteTalk first frame handling
if (extra_latents is not None
and not multitalk_sampling
and transformer.multitalk_model_type=="InfiniteTalk"):
for entry in extra_latents:
add_index = entry["index"]
num_extra_frames = entry["samples"].shape[2]
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
# differential diffusion inpaint
if masks is not None:
if idx < len(timesteps) - 1:
noise_timestep = timesteps[idx+1]
image_latent = sample_scheduler.scale_noise(
original_image.to(device), torch.tensor([noise_timestep]), noise.to(device)
)
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):
noisy_latents = add_noise(ttm_reference_latents, noise, sample_scheduler.all_timesteps[idx + ttm_start_step + 1].to(noise.device)).to(latent)
latent = latent * (1 - motion_mask) + noisy_latents * motion_mask
else:
latent = latent * (1 - motion_mask) + ttm_reference_latents.to(latent) * motion_mask
if freeinit_args is not None:
current_latent = latent.clone()
if callback is not None:
if recammaster is not None or mocha_embeds is not None:
callback_latent = (latent_model_input[:, :orig_noise_len].to(device) - noise_pred[:, :orig_noise_len].to(device) * t.to(device) / 1000).detach()
#elif phantom_latents is not None:
# callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach()
elif humo_reference_count > 0:
callback_latent = (latent_model_input[:,:-humo_reference_count].to(device) - noise_pred[:,:-humo_reference_count].to(device) * t.to(device) / 1000).detach()
elif "rcm" in sample_scheduler.__class__.__name__.lower():
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device)).detach()
else:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach()
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
else:
pbar.update(1)
except Exception as e:
log.error(f"Error during sampling: {e}")
if force_offload:
+85 -3
View File
@@ -4,17 +4,99 @@ import logging
import math
from tqdm import tqdm
from pathlib import Path
import os
import gc
import types, collections
from comfy.utils import ProgressBar, copy_to_param, set_attr_param
from comfy.model_patcher import get_key_weight, string_to_seed
from comfy.lora import calculate_weight
from comfy.model_management import cast_to_device
from comfy.float import stochastic_rounding
from .custom_linear import remove_lora_from_module
import folder_paths
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
log = logging.getLogger(__name__)
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
try:
from .gguf.gguf import GGUFParameter
except:
pass
class MetaParameter(torch.nn.Parameter):
def __new__(cls, dtype, quant_type=None):
data = torch.empty(0, dtype=dtype)
self = torch.nn.Parameter(data, requires_grad=False)
self.quant_type = quant_type
return self
def offload_transformer(transformer, remove_lora=True):
transformer.teacache_state.clear_all()
transformer.magcache_state.clear_all()
transformer.easycache_state.clear_all()
if transformer.patched_linear:
for name, param in transformer.named_parameters():
if "loras" in name or "controlnet" in name:
continue
module = transformer
subnames = name.split('.')
for subname in subnames[:-1]:
module = getattr(module, subname)
attr_name = subnames[-1]
if param.data.is_floating_point():
meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False)
setattr(module, attr_name, meta_param)
elif isinstance(param.data, GGUFParameter):
quant_type = getattr(param, 'quant_type', None)
setattr(module, attr_name, MetaParameter(param.data.dtype, quant_type))
else:
pass
if remove_lora:
remove_lora_from_module(transformer)
else:
transformer.to(offload_device)
for block in transformer.blocks:
block.kv_cache = None
if transformer.audio_model is not None and hasattr(block, 'audio_block'):
block.audio_block = None
mm.soft_empty_cache()
gc.collect()
def init_blockswap(transformer, block_swap_args, model):
if not transformer.patched_linear:
if block_swap_args is not None:
for name, param in transformer.named_parameters():
if "block" not in name or "control_adapter" in name or "face" 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()
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)
def check_device_same(first_device, second_device):
if first_device.type != second_device.type:
return False
@@ -140,7 +222,7 @@ def patch_weight_to_device(self, key, device_to=None, inplace_update=False, back
self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(weight.to(device=self.offload_device, copy=inplace_update), inplace_update)
if device_to is not None:
temp_weight = cast_to_device(weight, device_to, torch.float32, copy=True)
temp_weight = mm.cast_to_device(weight, device_to, torch.float32, copy=True)
else:
temp_weight = weight.to(torch.float32, copy=True)
if convert_func is not None:
+1 -1
View File
@@ -42,7 +42,7 @@ def _apply_custom_sigmas(sample_scheduler, sigmas, device):
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
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, enhance_hf=False, **kwargs):
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, denoise_strength=1.0, sigmas=None, log_timesteps=False, enhance_hf=False, **kwargs):
timesteps = None
if sigmas is not None:
steps = len(sigmas) - 1