diff --git a/nodes.py b/nodes.py index 9d2e709..594b9c3 100644 --- a/nodes.py +++ b/nodes.py @@ -1476,6 +1476,7 @@ class WanVideoExperimentalArgs: "fresca_freq_cutoff": ("INT", {"default": 20, "min": 0, "max": 10000, "step": 1}), "use_tcfg": ("BOOLEAN", {"default": False, "tooltip": "https://arxiv.org/abs/2503.18137 TCFG: Tangential Damping Classifier-free Guidance. CFG artifacts reduction."}), "raag_alpha": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Alpha value for RAAG, 1.0 is default, 0.0 is disabled."}), + "bidirectional_sampling": ("BOOLEAN", {"default": False, "tooltip": "Enable bidirectional sampling, based on https://github.com/ff2416/WanFM"}) }, } @@ -1648,6 +1649,8 @@ class WanVideoSampler: first_sampler = (end_step != -1 or end_step >= steps) + noise_pred_flipped = None + if isinstance(cfg, list): if steps != len(cfg): log.info(f"Received {len(cfg)} cfg values, but only {steps} steps. Setting step count to match.") @@ -2130,10 +2133,7 @@ class WanVideoSampler: freqs = None transformer.rope_embedder.k = None transformer.rope_embedder.num_frames = None - if "comfy" in rope_function: - transformer.rope_embedder.k = riflex_freq_index - transformer.rope_embedder.num_frames = latent_video_length - else: + if "default" in rope_function or bidirectional_sampling: d = transformer.dim // transformer.num_heads freqs = torch.cat([ rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index), @@ -2141,6 +2141,10 @@ class WanVideoSampler: rope_params(1024, 2 * (d // 6)) ], dim=1) + elif "comfy" in rope_function: + transformer.rope_embedder.k = riflex_freq_index + transformer.rope_embedder.num_frames = latent_video_length + transformer.rope_func = rope_function for block in transformer.blocks: block.rope_func = rope_function @@ -2256,7 +2260,7 @@ class WanVideoSampler: timesteps[-drift_steps:] = drift_timesteps[-drift_steps:] # Experimental args - use_cfg_zero_star = use_tangential = use_fresca = False + use_cfg_zero_star = use_tangential = use_fresca = bidirectional_sampling =False raag_alpha = 0.0 if experimental_args is not None: video_attention_split_steps = experimental_args.get("video_attention_split_steps", []) @@ -2277,10 +2281,15 @@ class WanVideoSampler: fresca_scale_high = experimental_args.get("fresca_scale_high", 1.25) fresca_freq_cutoff = experimental_args.get("fresca_freq_cutoff", 20) + bidirectional_sampling = experimental_args.get("bidirectional_sampling", False) + if bidirectional_sampling: + import copy + sample_scheduler_flipped = copy.deepcopy(sample_scheduler) + #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None, - add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None): + add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False): nonlocal transformer z = z.to(dtype) with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])): @@ -2324,8 +2333,14 @@ class WanVideoSampler: elif ATI_tracks is not None and ((ati_start_percent <= current_step_percentage <= ati_end_percent) or (ati_end_percent > 0 and idx == 0 and current_step_percentage >= ati_start_percent)): image_cond_input = image_cond_ati.to(z) - else: - image_cond_input = image_cond.to(z) if image_cond is not None else None + elif image_cond is not None: + if reverse_time: # Flip the image condition + image_cond_input = torch.cat([ + torch.flip(image_cond[:4], dims=[1]), + torch.flip(image_cond[4:], dims=[1]) + ]).to(z) + else: + image_cond_input = image_cond.to(z) if control_camera_latents is not None: if (control_camera_start_percent <= current_step_percentage <= control_camera_end_percent) or \ @@ -2442,6 +2457,7 @@ class WanVideoSampler: "inner_t": [shot_len] if shot_len else None, "standin_input": standin_input, "fantasy_portrait_input": fantasy_portrait_input, + "reverse_time": reverse_time } batch_size = 1 @@ -2701,6 +2717,10 @@ class WanVideoSampler: latent = image_latent * mask + latent * (1-mask) # end diff diff + if bidirectional_sampling: + latent_flipped = torch.flip(latent, dims=[1]) + latent_model_input_flipped = latent_flipped.to(device) + latent_model_input = latent.to(device) current_step_percentage = idx / len(timesteps) @@ -3237,6 +3257,14 @@ class WanVideoSampler: text_embeds["negative_prompt_embeds"], timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input) + if bidirectional_sampling: + noise_pred_flipped, self.cache_state = predict_with_cfg( + latent_model_input_flipped, + cfg[idx], + text_embeds["prompt_embeds"], + text_embeds["negative_prompt_embeds"], + timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, + cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, reverse_time=True) if latent_shift_loop: #reverse latent shift @@ -3276,7 +3304,15 @@ class WanVideoSampler: timestep, latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else 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 freeinit_args is not None: current_latent = latent.clone() diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 562478f..b17da61 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -31,7 +31,14 @@ from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot __all__ = ['WanModel'] from comfy import model_management as mm -from comfy.ldm.flux.math import apply_rope as apply_rope_comfy + +#from comfy.ldm.flux.math import apply_rope as apply_rope_comfy +def apply_rope_comfy(xq, xk, freqs_cis): + xq_ = xq.to(dtype=freqs_cis.dtype).reshape(*xq.shape[:-1], -1, 1, 2) + xk_ = xk.to(dtype=freqs_cis.dtype).reshape(*xk.shape[:-1], -1, 1, 2) + xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1] + xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1] + return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk) def apply_rope_comfy_chunked(xq, xk, freqs_cis, num_chunks=4): seq_dim = 1 @@ -159,7 +166,7 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0): @torch.autocast(device_type=mm.get_autocast_device(mm.get_torch_device()), enabled=False) @torch.compiler.disable() -def rope_apply(x, grid_sizes, freqs): +def rope_apply(x, grid_sizes, freqs, reverse_time=False): n, c = x.size(2), x.size(3) // 2 # split freqs @@ -173,12 +180,24 @@ def rope_apply(x, grid_sizes, freqs): # precompute multipliers x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape( seq_len, n, -1, 2)) - freqs_i = torch.cat([ - freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1), - freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1), - freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1) - ], - dim=-1).reshape(seq_len, 1, -1) + if reverse_time: + time_freqs = freqs[0][:f].view(f, 1, 1, -1) + time_freqs = torch.flip(time_freqs, dims=[0]) + time_freqs = time_freqs.expand(f, h, w, -1) + + spatial_freqs = torch.cat([ + freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1), + freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1) + ], dim=-1) + + freqs_i = torch.cat([time_freqs, spatial_freqs], dim=-1).reshape(seq_len, 1, -1) + else: + freqs_i = torch.cat([ + freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1), + freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1), + freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1) + ], + dim=-1).reshape(seq_len, 1, -1) # apply rotary embedding x_i = torch.view_as_real(x_i * freqs_i).flatten(2) @@ -760,6 +779,7 @@ class WanAttentionBlock(nn.Module): freqs_ip=None, adapter_proj=None, ip_scale=1.0, + reverse_time=False ): r""" Args: @@ -822,8 +842,8 @@ class WanAttentionBlock(nn.Module): elif self.rope_func == "comfy_chunked": q, k = apply_rope_comfy_chunked(q, k, freqs) else: - q=rope_apply(q, grid_sizes, freqs) - k=rope_apply(k, grid_sizes, freqs) + q=rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time) + k=rope_apply(k, grid_sizes, freqs, reverse_time=reverse_time) # FETA if enhance_enabled: @@ -1509,7 +1529,8 @@ class WanModel(torch.nn.Module): ref_target_masks=None, inner_t=None, standin_input=None, - fantasy_portrait_input=None + fantasy_portrait_input=None, + reverse_time=False ): r""" Forward pass through the diffusion model @@ -1652,7 +1673,8 @@ class WanModel(torch.nn.Module): if (self.cached_freqs is not None and self.cached_shape == current_shape and self.cached_cond == has_cond and - self.cached_rope_k == self.rope_embedder.k): + self.cached_rope_k == self.rope_embedder.k + ): freqs = self.cached_freqs else: img_ids = torch.zeros((f_len, h_len, w_len, 3), device=x.device, dtype=x.dtype) @@ -1965,7 +1987,8 @@ class WanModel(torch.nn.Module): freqs_ip=freqs_ip if x_ip is not None else None, e_ip=e0_ip if x_ip is not None else None, adapter_proj=adapter_proj, - ip_scale=ip_scale + ip_scale=ip_scale, + reverse_time=reverse_time ) if vace_data is not None: diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index a481673..5217225 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -6,7 +6,7 @@ from .flowmatch_pusa import FlowMatchSchedulerPusa from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep from .scheduling_flow_match_lcm import FlowMatchLCMScheduler from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, DEISMultistepScheduler - +import numpy as np from ...utils import log scheduler_list = [ @@ -14,7 +14,6 @@ scheduler_list = [ "dpm++", "dpm++/beta", "dpm++_sde", "dpm++_sde/beta", "euler", "euler/beta", - #"euler/accvideo", "deis", "lcm", "lcm/beta", "res_multistep", @@ -41,16 +40,7 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift)) else: sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) - # elif scheduler in ['euler/accvideo']: - # if steps != 50: - # raise Exception("Steps must be set to 50 for accvideo scheduler, 10 actual steps are used") - # sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta')) - # sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) - # start_latent_list = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50] - # sample_scheduler.sigmas = sample_scheduler.sigmas[start_latent_list] - # steps = len(start_latent_list) - 1 - # sample_scheduler.timesteps = timesteps = sample_scheduler.timesteps[start_latent_list[:steps]] - elif 'dpm++' in scheduler: + elif 'dpm' in scheduler: if 'sde' in scheduler: algorithm_type = "sde-dpmsolver++" else: