diff --git a/multitalk/multitalk_loop.py b/multitalk/multitalk_loop.py new file mode 100644 index 0000000..ed8b0af --- /dev/null +++ b/multitalk/multitalk_loop.py @@ -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}, diff --git a/nodes.py b/nodes.py index d754566..67f5142 100644 --- a/nodes.py +++ b/nodes.py @@ -1847,33 +1847,7 @@ 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", diff --git a/nodes_sampler.py b/nodes_sampler.py index 1adcc22..e1216a9 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -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: diff --git a/utils.py b/utils.py index ddfb1c1..243e5a6 100644 --- a/utils.py +++ b/utils.py @@ -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: diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index abb450f..b5d6c62 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -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