Merge branch 'main' into dev

This commit is contained in:
kijai
2025-08-17 13:52:30 +03:00
4 changed files with 85 additions and 36 deletions
+1 -1
View File
@@ -84,7 +84,7 @@ class MultiTalkWav2VecEmbeds:
"wav2vec_model": ("WAV2VECMODEL",),
"audio_1": ("AUDIO",),
"normalize_loudness": ("BOOLEAN", {"default": True}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 1000, "step": 1}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 1}),
"fps": ("FLOAT", {"default": 25.0, "min": 1.0, "max": 60.0, "step": 0.1}),
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "Strength of the audio conditioning"}),
"audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done: slower inference but more motion is allowed"}),
+45 -9
View File
@@ -1490,6 +1490,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"})
},
}
@@ -1679,6 +1680,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.")
@@ -2159,10 +2162,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),
@@ -2170,6 +2170,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
@@ -2282,7 +2286,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", [])
@@ -2303,10 +2307,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"])):
@@ -2350,8 +2359,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 \
@@ -2465,6 +2480,7 @@ class WanVideoSampler:
"standin_input": standin_input,
"fantasy_portrait_input": fantasy_portrait_input,
"phantom_ref": phantom_ref
"reverse_time": reverse_time
}
batch_size = 1
@@ -2724,6 +2740,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)
@@ -3260,6 +3280,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
@@ -3299,7 +3327,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()
+37 -14
View File
@@ -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:
@@ -1522,7 +1542,8 @@ class WanModel(torch.nn.Module):
inner_t=None,
standin_input=None,
fantasy_portrait_input=None,
phantom_ref=None
phantom_ref=None,
reverse_time=False
):
r"""
Forward pass through the diffusion model
@@ -1674,7 +1695,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)
@@ -1723,7 +1745,7 @@ class WanModel(torch.nn.Module):
ip_img_ids[:, :, :, 2] = ip_img_ids[:, :, :, 2] + torch.linspace(w_len + freq_offset, w_len + freq_offset + w_ip - 1, steps=w_ip, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
ip_img_ids = repeat(ip_img_ids, "t h w c -> b (t h w) c", b=1)
freqs_ip = self.rope_embedder(ip_img_ids).movedim(1, 2)
print("freqs_ip shape:", freqs_ip.shape)
#print("freqs_ip shape:", freqs_ip.shape)
# EchoShot cross attn freqs
inner_c = None
@@ -1987,7 +2009,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:
+2 -12
View File
@@ -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: