Merge branch 'main' into dev
This commit is contained in:
+1
-1
@@ -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"}),
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user