diff --git a/freeinit/freeinit_utils.py b/freeinit/freeinit_utils.py new file mode 100644 index 0000000..de7de14 --- /dev/null +++ b/freeinit/freeinit_utils.py @@ -0,0 +1,142 @@ +#https://github.com/TianxingWu/FreeInit/blob/master/freeinit_utils.py + +import torch +import torch.fft as fft +import math + + +def freq_mix_3d(x, noise, LPF): + """ + Noise reinitialization. + + Args: + x: diffused latent + noise: randomly sampled noise + LPF: low pass filter + """ + # FFT + x_freq = fft.fftn(x, dim=(-3, -2, -1)) + x_freq = fft.fftshift(x_freq, dim=(-3, -2, -1)) + noise_freq = fft.fftn(noise, dim=(-3, -2, -1)) + noise_freq = fft.fftshift(noise_freq, dim=(-3, -2, -1)) + + # frequency mix + HPF = 1 - LPF + x_freq_low = x_freq * LPF + noise_freq_high = noise_freq * HPF + x_freq_mixed = x_freq_low + noise_freq_high # mix in freq domain + + # IFFT + x_freq_mixed = fft.ifftshift(x_freq_mixed, dim=(-3, -2, -1)) + x_mixed = fft.ifftn(x_freq_mixed, dim=(-3, -2, -1)).real + + return x_mixed + + +def get_freq_filter(shape, device, filter_type, n, d_s, d_t): + """ + Form the frequency filter for noise reinitialization. + + Args: + shape: shape of latent (B, C, T, H, W) + filter_type: type of the freq filter + n: (only for butterworth) order of the filter, larger n ~ ideal, smaller n ~ gaussian + d_s: normalized stop frequency for spatial dimensions (0.0-1.0) + d_t: normalized stop frequency for temporal dimension (0.0-1.0) + """ + if filter_type == "gaussian": + return gaussian_low_pass_filter(shape=shape, d_s=d_s, d_t=d_t).to(device) + elif filter_type == "ideal": + return ideal_low_pass_filter(shape=shape, d_s=d_s, d_t=d_t).to(device) + elif filter_type == "box": + return box_low_pass_filter(shape=shape, d_s=d_s, d_t=d_t).to(device) + elif filter_type == "butterworth": + return butterworth_low_pass_filter(shape=shape, n=n, d_s=d_s, d_t=d_t).to(device) + else: + raise NotImplementedError + +def gaussian_low_pass_filter(shape, d_s=0.25, d_t=0.25): + """ + Compute the gaussian low pass filter mask. + + Args: + shape: shape of the filter (volume) + d_s: normalized stop frequency for spatial dimensions (0.0-1.0) + d_t: normalized stop frequency for temporal dimension (0.0-1.0) + """ + T, H, W = shape[-3], shape[-2], shape[-1] + mask = torch.zeros(shape) + if d_s==0 or d_t==0: + return mask + for t in range(T): + for h in range(H): + for w in range(W): + d_square = (((d_s/d_t)*(2*t/T-1))**2 + (2*h/H-1)**2 + (2*w/W-1)**2) + mask[..., t,h,w] = math.exp(-1/(2*d_s**2) * d_square) + return mask + + +def butterworth_low_pass_filter(shape, n=4, d_s=0.25, d_t=0.25): + """ + Compute the butterworth low pass filter mask. + + Args: + shape: shape of the filter (volume) + n: order of the filter, larger n ~ ideal, smaller n ~ gaussian + d_s: normalized stop frequency for spatial dimensions (0.0-1.0) + d_t: normalized stop frequency for temporal dimension (0.0-1.0) + """ + T, H, W = shape[-3], shape[-2], shape[-1] + mask = torch.zeros(shape) + if d_s==0 or d_t==0: + return mask + for t in range(T): + for h in range(H): + for w in range(W): + d_square = (((d_s/d_t)*(2*t/T-1))**2 + (2*h/H-1)**2 + (2*w/W-1)**2) + mask[..., t,h,w] = 1 / (1 + (d_square / d_s**2)**n) + return mask + + +def ideal_low_pass_filter(shape, d_s=0.25, d_t=0.25): + """ + Compute the ideal low pass filter mask. + + Args: + shape: shape of the filter (volume) + d_s: normalized stop frequency for spatial dimensions (0.0-1.0) + d_t: normalized stop frequency for temporal dimension (0.0-1.0) + """ + T, H, W = shape[-3], shape[-2], shape[-1] + mask = torch.zeros(shape) + if d_s==0 or d_t==0: + return mask + for t in range(T): + for h in range(H): + for w in range(W): + d_square = (((d_s/d_t)*(2*t/T-1))**2 + (2*h/H-1)**2 + (2*w/W-1)**2) + mask[..., t,h,w] = 1 if d_square <= d_s*2 else 0 + return mask + + +def box_low_pass_filter(shape, d_s=0.25, d_t=0.25): + """ + Compute the ideal low pass filter mask (approximated version). + + Args: + shape: shape of the filter (volume) + d_s: normalized stop frequency for spatial dimensions (0.0-1.0) + d_t: normalized stop frequency for temporal dimension (0.0-1.0) + """ + T, H, W = shape[-3], shape[-2], shape[-1] + mask = torch.zeros(shape) + if d_s==0 or d_t==0: + return mask + + threshold_s = round(int(H // 2) * d_s) + threshold_t = round(T // 2 * d_t) + + cframe, crow, ccol = T // 2, H // 2, W //2 + mask[..., cframe - threshold_t:cframe + threshold_t, crow - threshold_s:crow + threshold_s, ccol - threshold_s:ccol + threshold_s] = 1.0 + + return mask \ No newline at end of file diff --git a/nodes.py b/nodes.py index 1bb76b0..38fb6e6 100644 --- a/nodes.py +++ b/nodes.py @@ -1723,6 +1723,28 @@ class WanVideoExperimentalArgs: def process(self, **kwargs): return (kwargs,) +class WanVideoFreeInitArgs: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "freeinit_num_iters": ("INT", {"default": 3, "min": 1, "max": 10, "tooltip": "Number of FreeInit iterations"}), + "freeinit_method": (["butterworth", "ideal", "gaussian", "none"], {"default": "ideal", "tooltip": "Frequency filter type"}), + "freeinit_n": ("INT", {"default": 4, "min": 1, "max": 10, "tooltip": "Butterworth filter order (only for butterworth)"}), + "freeinit_d_s": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Spatial filter cutoff"}), + "freeinit_d_t": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Temporal filter cutoff"}), + }, + } + + RETURN_TYPES = ("FREEINITARGS", ) + RETURN_NAMES = ("freeinit_args",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "https://github.com/TianxingWu/FreeInit; FreeInit, a concise yet effective method to improve temporal consistency of videos generated by diffusion models" + EXPERIMENTAL = True + + def process(self, **kwargs): + return (kwargs,) + #region Sampler class WanVideoSampler: @classmethod @@ -1763,6 +1785,7 @@ class WanVideoSampler: "fantasytalking_embeds": ("FANTASYTALKING_EMBEDS", ), "uni3c_embeds": ("UNI3C_EMBEDS", ), "multitalk_embeds": ("MULTITALK_EMBEDS", ), + "freeinit_args": ("FREEINITARGS", ), } } @@ -1774,7 +1797,7 @@ class WanVideoSampler: def process(self, model, image_embeds, shift, steps, cfg, seed, scheduler, riflex_freq_index, text_embeds=None, 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): + experimental_args=None, sigmas=None, unianimate_poses=None, fantasytalking_embeds=None, uni3c_embeds=None, multitalk_embeds=None, freeinit_args=None): patcher = model model = model.model @@ -1877,7 +1900,7 @@ class WanVideoSampler: sample_scheduler.timesteps = denoising_step_list[:steps].clone().detach().to(device) sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)]) return sample_scheduler, timesteps - + if scheduler != "multitalk": sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, sigmas=sigmas) if timesteps is None: @@ -2062,6 +2085,24 @@ class WanVideoSampler: phantom_latents = phantom_latents.to(device) latent_video_length = noise.shape[1] + + # Initialize FreeInit filter if enabled + freq_filter = None + if freeinit_args is not None: + from .freeinit.freeinit_utils import get_freq_filter, freq_mix_3d + filter_shape = list(noise.shape) # [batch, C, T, H, W] + freq_filter = get_freq_filter( + filter_shape, + device=device, + filter_type=freeinit_args.get("freeinit_method", "butterworth"), + n=freeinit_args.get("freeinit_n", 4) if freeinit_args.get("freeinit_method", "butterworth") == "butterworth" else None, + d_s=freeinit_args.get("freeinit_s", 1.0), + d_t=freeinit_args.get("freeinit_t", 1.0) + ) + if samples is not None: + saved_generator_state = samples.get("generator_state", None) + if saved_generator_state is not None: + seed_g.set_state(saved_generator_state) if unianimate_poses is not None: transformer.dwpose_embedding.to(device, model["dtype"]) @@ -2724,56 +2765,156 @@ class WanVideoSampler: except: pass - #region main loop start - for idx, t in enumerate(tqdm(timesteps)): - if flowedit_args is not None: - if idx < skip_steps: + # Main sampling loop with FreeInit iterations + iterations = freeinit_args.get("freeinit_num_iters", 3) if freeinit_args is not None else 1 + current_latent = latent + + for iter_idx in range(iterations): + # 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, shift, device, sigmas=sigmas) + if timesteps is None: + timesteps = sample_scheduler.timesteps + + # Diffuse current latent to t=999 + diffuse_timesteps = torch.full((noise.shape[0],), 999, device=device, dtype=torch.long) + z_T = add_noise( + current_latent.to(device), + initial_noise_saved.to(device), + diffuse_timesteps + ) + + # Generate new random noise + z_rand = torch.randn(z_T.shape, dtype=torch.float32, generator=seed_g, device=torch.device("cpu")) + + # Apply frequency mixing + current_latent = freq_mix_3d(z_T.to(torch.float32), z_rand.to(device), LPF=freq_filter) + current_latent = current_latent.to(dtype) + + # Store initial noise for first iteration + if iter_idx == 0: + initial_noise_saved = current_latent.detach().clone() + if samples is not None: + current_latent = input_samples.to(device) continue + + # Reset per-iteration states + self.cache_state = [None, None] + self.cache_state_source = [None, None] + self.cache_states_context = [] + if context_options is not None: + self.window_tracker = WindowTracker(verbose=context_options["verbose"]) + + # Set latent for denoising + latent = current_latent - # diff diff - if masks is not None: - if idx < len(timesteps) - 1: - noise_timestep = timesteps[idx+1] - image_latent = sample_scheduler.scale_noise( - original_image, torch.tensor([noise_timestep]), noise.to(device) - ) - mask = masks[idx] - mask = mask.to(latent) - latent = image_latent * mask + latent * (1-mask) - # end diff diff + print(latent) - latent_model_input = latent.to(device) + #region main loop start + for idx, t in enumerate(tqdm(timesteps)): + if flowedit_args is not None: + if idx < skip_steps: + continue - timestep = torch.tensor([t]).to(device) - current_step_percentage = idx / len(timesteps) + # diff diff + if masks is not None: + if idx < len(timesteps) - 1: + noise_timestep = timesteps[idx+1] + image_latent = sample_scheduler.scale_noise( + original_image, torch.tensor([noise_timestep]), noise.to(device) + ) + mask = masks[idx] + mask = mask.to(latent) + latent = image_latent * mask + latent * (1-mask) + # end diff diff - ### latent shift - if latent_shift_loop: - if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent: - latent_model_input = torch.cat([latent_model_input[:, shift_idx:]] + [latent_model_input[:, :shift_idx]], dim=1) + latent_model_input = latent.to(device) - #enhance-a-video - if feta_args is not None and feta_start_percent <= current_step_percentage <= feta_end_percent: - enable_enhance() - else: - disable_enhance() + timestep = torch.tensor([t]).to(device) + current_step_percentage = idx / len(timesteps) - #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 + ### latent shift + if latent_shift_loop: + if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent: + latent_model_input = torch.cat([latent_model_input[:, shift_idx:]] + [latent_model_input[:, :shift_idx]], dim=1) - #source - if idx < len(timesteps) - drift_steps: + #enhance-a-video + if feta_args is not None and feta_start_percent <= current_step_percentage <= feta_end_percent: + enable_enhance() + else: + disable_enhance() + + #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_src, device=intermediate_device) - vt_src = torch.zeros_like(zt_src, device=intermediate_device) + 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) @@ -2786,52 +2927,58 @@ class WanVideoSampler: 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] + + if len(text_embeds["prompt_embeds"]) > 1: + positive = text_embeds["prompt_embeds"][prompt_index] else: - positive = source_embeds["prompt_embeds"] - + positive = text_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_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_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) + 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_src_context, c, latent_video_length, context_overlap) - vt_src[:, c, :, :] += vt_src_context * window_mask + + 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_src /= counter + vt_tgt /= 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) + 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 + #context windowing + elif context_options is not None: + counter = torch.zeros_like(latent_model_input, device=intermediate_device) + noise_pred = torch.zeros_like(latent_model_input, 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) + fraction_per_context = 1.0 / len(context_queue) + context_pbar = ProgressBar(steps) + step_start_progress = idx + for i, c in enumerate(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: @@ -2840,451 +2987,397 @@ class WanVideoSampler: prompt_index = min(int(max(c) / section_size), num_prompts - 1) if context_options["verbose"]: log.info(f"Prompt index: {prompt_index}") - + + # Use the appropriate prompt for this section 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_img_emb = image_cond[:, c] + partial_img_emb[:, 0] = image_cond[:, 0].to(intermediate_device) - 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 control_latents is not None: + partial_control_latents = control_latents[:, c] + partial_control_camera_latents = None + if control_camera_latents is not None: + partial_control_camera_latents = control_camera_latents[:, :, c] + + partial_vace_context = None + if vace_data is not None: + window_vace_data = [] + for vace_entry in vace_data: + partial_context = vace_entry["context"][0][:, c] + if has_ref: + partial_context[:, 0] = vace_entry["context"][0][:, 0] + + window_vace_data.append({ + "context": [partial_context], + "scale": vace_entry["scale"], + "start": vace_entry["start"], + "end": vace_entry["end"], + "seq_len": vace_entry["seq_len"] + }) + + partial_vace_context = window_vace_data + + partial_audio_proj = None + if fantasytalking_embeds is not None: + partial_audio_proj = audio_proj[:, c] + + partial_latent_model_input = latent_model_input[:, c] + + partial_unianim_data = None + if unianim_data is not None: + partial_dwpose = dwpose_data[:, :, c] + partial_dwpose_flat=rearrange(partial_dwpose, 'b c f h w -> b (f h w) c') + partial_unianim_data = { + "dwpose": partial_dwpose_flat, + "random_ref": unianim_data["random_ref"], + "strength": unianimate_poses["strength"], + "start_percent": unianimate_poses["start_percent"], + "end_percent": unianimate_poses["end_percent"] + } + + partial_add_cond = None + if add_cond is not None: + partial_add_cond = add_cond[:, :, c].to(device, dtype) + + noise_pred_context, new_teacache = predict_with_cfg( + partial_latent_model_input, + cfg[idx], positive, + text_embeds["negative_prompt_embeds"], + timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj, + partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c) + 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 - #context windowing - elif context_options is not None: - counter = torch.zeros_like(latent_model_input, device=intermediate_device) - noise_pred = torch.zeros_like(latent_model_input, device=intermediate_device) - context_queue = list(context(idx, steps, latent_video_length, context_frames, context_stride, context_overlap)) - fraction_per_context = 1.0 / len(context_queue) - context_pbar = ProgressBar(steps) - step_start_progress = idx - for i, c in enumerate(context_queue): - window_id = self.window_tracker.get_window_id(c) + window_mask = create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=is_looped) + noise_pred[:, c] += noise_pred_context * window_mask + counter[:, c] += window_mask + context_pbar.update_absolute(step_start_progress + (i + 1) * fraction_per_context, steps) + noise_pred /= counter + #region multitalk + elif multitalk_sampling: + original_image = cond_image = image_embeds.get("multitalk_start_image", None) + offload = image_embeds.get("force_offload", False) + tiled_vae = image_embeds.get("tiled_vae", False) + frame_num = clip_length = image_embeds.get("num_frames", 81) + vae = image_embeds.get("vae", None) + clip_embeds = image_embeds.get("clip_context", None) + 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) + + gen_video_list = [] + is_first_clip = True + arrive_last_frame = False + cur_motion_frames_num = 1 + audio_start_idx = iteration_count = 0 + audio_end_idx = audio_start_idx + clip_length + indices = (torch.arange(4 + 1) - 2) * 1 - 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}") - - # Use the appropriate prompt for this section - 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_control_camera_latents = None - if control_camera_latents is not None: - partial_control_camera_latents = control_camera_latents[:, :, c] - - partial_vace_context = None - if vace_data is not None: - window_vace_data = [] - for vace_entry in vace_data: - partial_context = vace_entry["context"][0][:, c] - if has_ref: - partial_context[:, 0] = vace_entry["context"][0][:, 0] - - window_vace_data.append({ - "context": [partial_context], - "scale": vace_entry["scale"], - "start": vace_entry["start"], - "end": vace_entry["end"], - "seq_len": vace_entry["seq_len"] - }) - - partial_vace_context = window_vace_data - - partial_audio_proj = None - if fantasytalking_embeds is not None: - partial_audio_proj = audio_proj[:, c] - - partial_latent_model_input = latent_model_input[:, c] - - partial_unianim_data = None - if unianim_data is not None: - partial_dwpose = dwpose_data[:, :, c] - partial_dwpose_flat=rearrange(partial_dwpose, 'b c f h w -> b (f h w) c') - partial_unianim_data = { - "dwpose": partial_dwpose_flat, - "random_ref": unianim_data["random_ref"], - "strength": unianimate_poses["strength"], - "start_percent": unianimate_poses["start_percent"], - "end_percent": unianimate_poses["end_percent"] - } - - partial_add_cond = None - if add_cond is not None: - partial_add_cond = add_cond[:, :, c].to(device, dtype) - - noise_pred_context, new_teacache = predict_with_cfg( - partial_latent_model_input, - cfg[idx], positive, - text_embeds["negative_prompt_embeds"], - timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj, - partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c) - - if cache_args is not None: - self.window_tracker.cache_states[window_id] = new_teacache - - window_mask = create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=is_looped) - noise_pred[:, c] += noise_pred_context * window_mask - counter[:, c] += window_mask - context_pbar.update_absolute(step_start_progress + (i + 1) * fraction_per_context, steps) - noise_pred /= counter - #region multitalk - elif multitalk_sampling: - original_image = cond_image = image_embeds.get("multitalk_start_image", None) - offload = image_embeds.get("force_offload", False) - tiled_vae = image_embeds.get("tiled_vae", False) - frame_num = clip_length = image_embeds.get("num_frames", 81) - vae = image_embeds.get("vae", None) - clip_embeds = image_embeds.get("clip_context", None) - 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) - - gen_video_list = [] - is_first_clip = True - arrive_last_frame = False - cur_motion_frames_num = 1 - audio_start_idx = iteration_count = 0 - audio_end_idx = audio_start_idx + clip_length - indices = (torch.arange(4 + 1) - 2) * 1 - - if multitalk_embeds is not None: - total_frames = len(multitalk_audio_embedding) - - estimated_iterations = total_frames // (frame_num - motion_frame) + 1 - loop_pbar = tqdm(total=estimated_iterations, desc="Generating video clips") - callback = prepare_callback(patcher, estimated_iterations) - - audio_embedding = multitalk_audio_embedding - human_num = len(audio_embedding) - audio_embs = None - while True: # start video generation iteratively 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, 1).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) + total_frames = len(multitalk_audio_embedding) - h, w = cond_image.shape[-2], cond_image.shape[-1] - 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]) + estimated_iterations = total_frames // (frame_num - motion_frame) + 1 + loop_pbar = tqdm(total=estimated_iterations, desc="Generating video clips") + callback = prepare_callback(patcher, estimated_iterations) - noise = torch.randn( - 16, (frame_num - 1) // 4 + 1, - lat_h, lat_w, dtype=torch.float32, device=device) + audio_embedding = multitalk_audio_embedding + human_num = len(audio_embedding) + audio_embs = None + while True: # start video generation iteratively + 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, 1).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) - # get mask - msk = torch.ones(1, frame_num, lat_h, lat_w, device=device) - msk[:, cur_motion_frames_num:] = 0 - msk = torch.concat([ - torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:] - ], dim=1) - msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w) - msk = msk.transpose(1, 2).to(dtype) # B 4 T H W + h, w = cond_image.shape[-2], cond_image.shape[-1] + 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]) - mm.soft_empty_cache() + noise = torch.randn( + 16, (frame_num - 1) // 4 + 1, + lat_h, lat_w, dtype=torch.float32, device=device) - # zero padding and vae encode - video_frames = torch.zeros(1, cond_image.shape[1], frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype) - padding_frames_pixels_values = torch.concat([cond_image.to(device, vae.dtype), video_frames], dim=2) + # get mask + msk = torch.ones(1, frame_num, lat_h, lat_w, device=device) + msk[:, cur_motion_frames_num:] = 0 + msk = torch.concat([ + torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:] + ], dim=1) + msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w) + msk = msk.transpose(1, 2).to(dtype) # B 4 T H W - vae.to(device) - y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae).to(dtype) - vae.to(offload_device) + mm.soft_empty_cache() - cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4) - latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W - y = torch.concat([msk, y], dim=1) # B 4+C T H W - mm.soft_empty_cache() + # zero padding and vae encode + video_frames = torch.zeros(1, cond_image.shape[1], frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype) + padding_frames_pixels_values = torch.concat([cond_image.to(device, vae.dtype), video_frames], dim=2) - 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: - sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, sigmas=sigmas) - if timesteps is None: - timesteps = sample_scheduler.timesteps + vae.to(device) + y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae).to(dtype) + vae.to(offload_device) - transformed_timesteps = [] - for t in timesteps: - t_tensor = torch.tensor([t.item()], device=device) - transformed_timesteps.append(t_tensor) + cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4) + latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W + y = torch.concat([msk, y], dim=1) # B 4+C T H W + mm.soft_empty_cache() - transformed_timesteps.append(torch.tensor([0.], device=device)) - timesteps = transformed_timesteps - - # sample videos - latent = noise - - # injecting motion frames - if not is_first_clip: - latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device) - motion_add_noise = torch.randn_like(latent_motion_frames).contiguous() - add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0]) - _, T_m, _, _ = add_latent.shape - latent[:, :T_m] = add_latent - - if offload: - #blockswap init - if transformer_options is not None: - block_swap_args = transformer_options.get("block_swap_args", None) - - if block_swap_args is not None: - transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True) - for name, param in transformer.named_parameters(): - if "block" not in name: - param.data = param.data.to(device) - if "control_adapter" 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, non_blocking=transformer.use_non_blocking) - elif block_swap_args["offload_img_emb"] and "img_emb" in name: - param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking) - - 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() - elif model["manual_offloading"]: - transformer.to(device) - - comfy_pbar = ProgressBar(len(timesteps)-1) - for i in tqdm(range(len(timesteps)-1)): - timestep = timesteps[i] - latent_model_input = latent.to(device) - - noise_pred, self.cache_state = predict_with_cfg( - latent_model_input, - cfg[idx], - text_embeds["prompt_embeds"], - text_embeds["negative_prompt_embeds"], - timestep, idx, y.squeeze(0), clip_embeds.to(dtype), control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, - cache_state=self.cache_state, multitalk_audio_embeds=audio_embs) - - 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(iteration_count, callback_latent, None, estimated_iterations) - - # update latent if scheduler == "multitalk": - noise_pred = -noise_pred - dt = timesteps[i] - timesteps[i + 1] - dt = dt / 1000 - latent = latent + noise_pred * dt[:, None, None, None] + 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: - latent = latent.to(intermediate_device) - step_args = { - "generator": seed_g, - } - if isinstance(sample_scheduler, DEISMultistepScheduler) or isinstance(sample_scheduler, FlowMatchScheduler): - step_args.pop("generator", None) - temp_x0 = sample_scheduler.step( - noise_pred.unsqueeze(0), - timestep, - latent.unsqueeze(0), - #return_dict=False, - **step_args)[0] - latent = temp_x0.squeeze(0) + sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, sigmas=sigmas) + if timesteps is None: + timesteps = sample_scheduler.timesteps + + transformed_timesteps = [] + for t in timesteps: + t_tensor = torch.tensor([t.item()], device=device) + transformed_timesteps.append(t_tensor) + + transformed_timesteps.append(torch.tensor([0.], device=device)) + timesteps = transformed_timesteps + + # sample videos + latent = noise # injecting motion frames if not is_first_clip: latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device) motion_add_noise = torch.randn_like(latent_motion_frames).contiguous() - add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1]) + add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0]) _, T_m, _, _ = add_latent.shape latent[:, :T_m] = add_latent - x0 = latent.to(device) - del latent_model_input, timestep - comfy_pbar.update(1) + if offload: + #blockswap init + if transformer_options is not None: + block_swap_args = transformer_options.get("block_swap_args", None) - if offload: - transformer.to(offload_device) - vae.to(device) - videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae) - vae.to(offload_device) + if block_swap_args is not None: + transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True) + for name, param in transformer.named_parameters(): + if "block" not in name: + param.data = param.data.to(device) + if "control_adapter" 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, non_blocking=transformer.use_non_blocking) + elif block_swap_args["offload_img_emb"] and "img_emb" in name: + param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking) + + 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() + elif model["manual_offloading"]: + transformer.to(device) + + comfy_pbar = ProgressBar(len(timesteps)-1) + for i in tqdm(range(len(timesteps)-1)): + timestep = timesteps[i] + latent_model_input = latent.to(device) + + noise_pred, self.cache_state = predict_with_cfg( + latent_model_input, + cfg[idx], + text_embeds["prompt_embeds"], + text_embeds["negative_prompt_embeds"], + timestep, idx, y.squeeze(0), clip_embeds.to(dtype), control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, + cache_state=self.cache_state, multitalk_audio_embeds=audio_embs) + + 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(iteration_count, callback_latent, None, estimated_iterations) + + # update latent + if scheduler == "multitalk": + noise_pred = -noise_pred + dt = timesteps[i] - timesteps[i + 1] + dt = dt / 1000 + latent = latent + noise_pred * dt[:, None, None, None] + else: + latent = latent.to(intermediate_device) + step_args = { + "generator": seed_g, + } + if isinstance(sample_scheduler, DEISMultistepScheduler) or isinstance(sample_scheduler, FlowMatchScheduler): + step_args.pop("generator", None) + temp_x0 = sample_scheduler.step( + noise_pred.unsqueeze(0), + timestep, + latent.unsqueeze(0), + #return_dict=False, + **step_args)[0] + latent = temp_x0.squeeze(0) + + # injecting motion frames + if not is_first_clip: + latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device) + motion_add_noise = torch.randn_like(latent_motion_frames).contiguous() + add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1]) + _, T_m, _, _ = add_latent.shape + latent[:, :T_m] = add_latent + + x0 = latent.to(device) + del latent_model_input, timestep + comfy_pbar.update(1) + + if offload: + transformer.to(offload_device) + vae.to(device) + videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae) + vae.to(offload_device) + + # cache generated samples + videos = torch.stack(videos).cpu() # B C T H W + if colormatch != "disabled": + videos = videos[0].permute(1, 2, 3, 0).cpu().numpy() + from color_matcher import ColorMatcher + cm = ColorMatcher() + cm_result_list = [] + for img in videos: + cm_result = cm.transfer(src=img, ref=original_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch) + cm_result_list.append(torch.from_numpy(cm_result)) - # cache generated samples - videos = torch.stack(videos).cpu() # B C T H W - if colormatch != "disabled": - videos = videos[0].permute(1, 2, 3, 0).cpu().numpy() - from color_matcher import ColorMatcher - cm = ColorMatcher() - cm_result_list = [] - for img in videos: - cm_result = cm.transfer(src=img, ref=original_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch) - cm_result_list.append(torch.from_numpy(cm_result)) + videos = torch.stack(cm_result_list, dim=0).to(torch.float32).permute(3, 0, 1, 2).unsqueeze(0) + + if is_first_clip: + gen_video_list.append(videos) + else: + gen_video_list.append(videos[:, :, cur_motion_frames_num:]) + + # decide whether is done + if arrive_last_frame: + loop_pbar.update(estimated_iterations - iteration_count) + loop_pbar.close() + break + + # update next condition frames + is_first_clip = False + cur_motion_frames_num = motion_frame + + cond_image = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device) + + # Update progress bar + iteration_count += 1 + loop_pbar.update(1) + + # Repeat audio emb + if multitalk_embeds is not None: + audio_start_idx += (frame_num - cur_motion_frames_num) + 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(1): + source_frame = len(audio_embedding[human_inx]) + source_frames.append(source_frame) + if audio_end_idx >= len(audio_embedding[human_inx]): + miss_length = audio_end_idx - len(audio_embedding[human_inx]) + 3 + 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], dim=0) + miss_lengths.append(miss_length) + else: + miss_lengths.append(0) + + gen_video_samples = torch.cat(gen_video_list, dim=2).to(torch.float32) + + del noise, latent + if force_offload: + if model["manual_offloading"]: + transformer.to(offload_device) + mm.soft_empty_cache() + gc.collect() + try: + print_memory(device) + torch.cuda.reset_peak_memory_stats(device) + except: + pass + return {"video": gen_video_samples[0].permute(1, 2, 3, 0).cpu()}, - videos = torch.stack(cm_result_list, dim=0).to(torch.float32).permute(3, 0, 1, 2).unsqueeze(0) - - if is_first_clip: - gen_video_list.append(videos) - else: - gen_video_list.append(videos[:, :, cur_motion_frames_num:]) - - # decide whether is done - if arrive_last_frame: - loop_pbar.update(estimated_iterations - iteration_count) - loop_pbar.close() - break - - # update next condition frames - is_first_clip = False - cur_motion_frames_num = motion_frame - - cond_image = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device) - - # Update progress bar - iteration_count += 1 - loop_pbar.update(1) - - # Repeat audio emb - if multitalk_embeds is not None: - audio_start_idx += (frame_num - cur_motion_frames_num) - 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(1): - source_frame = len(audio_embedding[human_inx]) - source_frames.append(source_frame) - if audio_end_idx >= len(audio_embedding[human_inx]): - miss_length = audio_end_idx - len(audio_embedding[human_inx]) + 3 - 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], dim=0) - miss_lengths.append(miss_length) - else: - miss_lengths.append(0) - - gen_video_samples = torch.cat(gen_video_list, dim=2).to(torch.float32) - - del noise, latent - if force_offload: - if model["manual_offloading"]: - transformer.to(offload_device) - mm.soft_empty_cache() - gc.collect() - try: - print_memory(device) - torch.cuda.reset_peak_memory_stats(device) - except: - pass - return {"video": gen_video_samples[0].permute(1, 2, 3, 0).cpu()}, - - #region normal inference - else: - noise_pred, self.cache_state = predict_with_cfg( - latent_model_input, - cfg[idx], - text_embeds["prompt_embeds"], - text_embeds["negative_prompt_embeds"], - timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, - cache_state=self.cache_state) - - if latent_shift_loop: - #reverse latent shift - if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent: - 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 - - - if flowedit_args is None: - latent = latent.to(intermediate_device) - step_args = { - "generator": seed_g, - } - if isinstance(sample_scheduler, DEISMultistepScheduler) or isinstance(sample_scheduler, FlowMatchScheduler): - step_args.pop("generator", None) - temp_x0 = sample_scheduler.step( - noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else noise_pred.unsqueeze(0), - t, - latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else latent.unsqueeze(0), - #return_dict=False, - **step_args)[0] - latent = temp_x0.squeeze(0) - - x0 = latent.to(device) - if callback is not None: - if recammaster 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().permute(1,0,2,3) - 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().permute(1,0,2,3) - else: - callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3) - callback(idx, callback_latent, None, steps) + #region normal inference else: - pbar.update(1) - del latent_model_input, timestep - else: - if callback is not None: - callback_latent = (zt_tgt.to(device) - vt_tgt.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3) - callback(idx, callback_latent, None, steps) + noise_pred, self.cache_state = predict_with_cfg( + latent_model_input, + cfg[idx], + text_embeds["prompt_embeds"], + text_embeds["negative_prompt_embeds"], + timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, + cache_state=self.cache_state) + + if latent_shift_loop: + #reverse latent shift + if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent: + 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 + + + if flowedit_args is None: + latent = latent.to(intermediate_device) + step_args = { + "generator": seed_g, + } + if isinstance(sample_scheduler, DEISMultistepScheduler) or isinstance(sample_scheduler, FlowMatchScheduler): + step_args.pop("generator", None) + temp_x0 = sample_scheduler.step( + noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else noise_pred.unsqueeze(0), + t, + latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else latent.unsqueeze(0), + #return_dict=False, + **step_args)[0] + latent = temp_x0.squeeze(0) + + x0 = latent.to(device) + + generator_state = seed_g.get_state() + + if freeinit_args is not None: + current_latent = x0.clone() + + if callback is not None: + if recammaster 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().permute(1,0,2,3) + 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().permute(1,0,2,3) + else: + callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3) + callback(idx, callback_latent, None, steps) + else: + pbar.update(1) + del latent_model_input, timestep else: - pbar.update(1) + if callback is not None: + callback_latent = (zt_tgt.to(device) - vt_tgt.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3) + callback(idx, callback_latent, None, steps) + else: + pbar.update(1) if phantom_latents is not None: x0 = x0[:,:-phantom_latents.shape[1]] @@ -3317,8 +3410,13 @@ class WanVideoSampler: pass return ({ - "samples": x0.unsqueeze(0).cpu(), "looped": is_looped, "end_image": end_image if not fun_or_fl2v_model else None, "has_ref": has_ref, "drop_last": drop_last, - }, ) + "samples": x0.unsqueeze(0).cpu(), + "looped": is_looped, + "end_image": end_image if not fun_or_fl2v_model else None, + "has_ref": has_ref, + "drop_last": drop_last, + "generator_state": generator_state, + }, ) class WindowTracker: def __init__(self, verbose=False): @@ -3561,6 +3659,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoRealisDanceLatents": WanVideoRealisDanceLatents, "WanVideoApplyNAG": WanVideoApplyNAG, "WanVideoMiniMaxRemoverEmbeds": WanVideoMiniMaxRemoverEmbeds, + "WanVideoFreeInitArgs": WanVideoFreeInitArgs, } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -3598,4 +3697,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoRealisDanceLatents": "WanVideo RealisDance Latents", "WanVideoApplyNAG": "WanVideo Apply NAG", "WanVideoMiniMaxRemoverEmbeds": "WanVideo MiniMax Remover Embeds", + "WanVideoFreeInitArgs": "WanVideo Free Init Args", }