Merge branch 'main' into lynx
This commit is contained in:
File diff suppressed because it is too large
Load Diff
+2
-1
@@ -13,9 +13,10 @@ def fp8_linear_forward(cls, base_dtype, input):
|
||||
if scale_weight is None:
|
||||
scale_weight = torch.ones((), device=input.device, dtype=torch.float32)
|
||||
else:
|
||||
scale_weight = scale_weight.to(input.device)
|
||||
scale_weight = scale_weight.to(input.device).squeeze()
|
||||
|
||||
scale_input = torch.ones((), device=input.device, dtype=torch.float32)
|
||||
|
||||
input = torch.clamp(input, min=-448, max=448, out=input)
|
||||
inn = input.reshape(-1, input_shape[2]).to(torch.float8_e4m3fn).contiguous() #always e4m3fn because e5m2 * e5m2 is not supported
|
||||
|
||||
|
||||
@@ -1738,7 +1738,10 @@ 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"})
|
||||
"bidirectional_sampling": ("BOOLEAN", {"default": False, "tooltip": "Enable bidirectional sampling, based on https://github.com/ff2416/WanFM"}),
|
||||
"temporal_score_rescaling": ("BOOLEAN", {"default": False, "tooltip": "Enable temporal score rescaling: https://github.com/temporalscorerescaling/TSR/"}),
|
||||
"tsr_k": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "The sampling temperature"}),
|
||||
"tsr_sigma": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "How early TSR steer the sampling process"}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -532,6 +532,8 @@ class WanVideoLoraSelectMulti:
|
||||
"low_mem_load": low_mem_load,
|
||||
"merge_loras": merge_loras,
|
||||
})
|
||||
if len(loras_list) == 0:
|
||||
return None,
|
||||
return (loras_list,)
|
||||
|
||||
class WanVideoVACEModelSelect:
|
||||
@@ -1220,6 +1222,15 @@ class WanVideoModelLoader:
|
||||
"e": [-114.36346466, 65.26524496, -18.82220707, 4.91518089, -0.23412683],
|
||||
"e0": [8.10705460e+03, 2.13393892e+03, -3.72934672e+02, 1.66203073e+01, -4.17769401e-02],
|
||||
},
|
||||
# Placeholders until TeaCache for Wan2.2 is obtained
|
||||
"14B_2.2": {
|
||||
"e": [-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429, -13.02252404],
|
||||
"e0": [-3.03318725e+05, 4.90537029e+04, -2.65530556e+03, 5.87365115e+01, -3.15583525e-01],
|
||||
},
|
||||
"i2v_14B_2.2":{
|
||||
"e": [-114.36346466, 65.26524496, -18.82220707, 4.91518089, -0.23412683],
|
||||
"e0": [8.10705460e+03, 2.13393892e+03, -3.72934672e+02, 1.66203073e+01, -4.17769401e-02],
|
||||
},
|
||||
}
|
||||
|
||||
magcache_ratios_map = {
|
||||
@@ -1227,6 +1238,8 @@ class WanVideoModelLoader:
|
||||
"14B": np.array([1.0]*2+[1.02504, 1.03017, 1.00025, 1.00251, 0.9985, 0.99962, 0.99779, 0.99771, 0.9966, 0.99658, 0.99482, 0.99476, 0.99467, 0.99451, 0.99664, 0.99656, 0.99434, 0.99431, 0.99533, 0.99545, 0.99468, 0.99465, 0.99438, 0.99434, 0.99516, 0.99517, 0.99384, 0.9938, 0.99404, 0.99401, 0.99517, 0.99516, 0.99409, 0.99408, 0.99428, 0.99426, 0.99347, 0.99343, 0.99418, 0.99416, 0.99271, 0.99269, 0.99313, 0.99311, 0.99215, 0.99215, 0.99218, 0.99215, 0.99216, 0.99217, 0.99163, 0.99161, 0.99138, 0.99135, 0.98982, 0.9898, 0.98996, 0.98995, 0.9887, 0.98866, 0.98772, 0.9877, 0.98767, 0.98765, 0.98573, 0.9857, 0.98501, 0.98498, 0.9838, 0.98376, 0.98177, 0.98173, 0.98037, 0.98035, 0.97678, 0.97677, 0.97546, 0.97543, 0.97184, 0.97183, 0.96711, 0.96708, 0.96349, 0.96345, 0.95629, 0.95625, 0.94926, 0.94929, 0.93964, 0.93961, 0.92511, 0.92504, 0.90693, 0.90678, 0.8796, 0.87945, 0.86111, 0.86189]),
|
||||
"i2v_480": np.array([1.0]*2+[0.98783, 0.98993, 0.97559, 0.97593, 0.98311, 0.98319, 0.98202, 0.98225, 0.9888, 0.98878, 0.98762, 0.98759, 0.98957, 0.98971, 0.99052, 0.99043, 0.99383, 0.99384, 0.98857, 0.9886, 0.99065, 0.99068, 0.98845, 0.98847, 0.99057, 0.99057, 0.98957, 0.98961, 0.98601, 0.9861, 0.98823, 0.98823, 0.98756, 0.98759, 0.98808, 0.98814, 0.98721, 0.98724, 0.98571, 0.98572, 0.98543, 0.98544, 0.98157, 0.98165, 0.98411, 0.98413, 0.97952, 0.97953, 0.98149, 0.9815, 0.9774, 0.97742, 0.97825, 0.97826, 0.97355, 0.97361, 0.97085, 0.97087, 0.97056, 0.97055, 0.96588, 0.96587, 0.96113, 0.96124, 0.9567, 0.95681, 0.94961, 0.94969, 0.93973, 0.93988, 0.93217, 0.93224, 0.91878, 0.91896, 0.90955, 0.90954, 0.92617, 0.92616]),
|
||||
"i2v_720": np.array([1.0]*2+[0.99428, 0.99498, 0.98588, 0.98621, 0.98273, 0.98281, 0.99018, 0.99023, 0.98911, 0.98917, 0.98646, 0.98652, 0.99454, 0.99456, 0.9891, 0.98909, 0.99124, 0.99127, 0.99102, 0.99103, 0.99215, 0.99212, 0.99515, 0.99515, 0.99576, 0.99572, 0.99068, 0.99072, 0.99097, 0.99097, 0.99166, 0.99169, 0.99041, 0.99042, 0.99201, 0.99198, 0.99101, 0.99101, 0.98599, 0.98603, 0.98845, 0.98844, 0.98848, 0.98851, 0.98862, 0.98857, 0.98718, 0.98719, 0.98497, 0.98497, 0.98264, 0.98263, 0.98389, 0.98393, 0.97938, 0.9794, 0.97535, 0.97536, 0.97498, 0.97499, 0.973, 0.97301, 0.96827, 0.96828, 0.96261, 0.96263, 0.95335, 0.9534, 0.94649, 0.94655, 0.93397, 0.93414, 0.91636, 0.9165, 0.89088, 0.89109, 0.8679, 0.86768]),
|
||||
"14B_2.2": np.array([1.0]*2+[0.99505, 0.99389, 0.99441, 0.9957, 0.99558, 0.99551, 0.99499, 0.9945, 0.99534, 0.99548, 0.99468, 0.9946, 0.99463, 0.99458, 0.9946, 0.99453, 0.99408, 0.99404, 0.9945, 0.99441, 0.99409, 0.99398, 0.99403, 0.99397, 0.99382, 0.99377, 0.99349, 0.99343, 0.99377, 0.99378, 0.9933, 0.99328, 0.99303, 0.99301, 0.99217, 0.99216, 0.992, 0.99201, 0.99201, 0.99202, 0.99133, 0.99132, 0.99112, 0.9911, 0.99155, 0.99155, 0.98958, 0.98957, 0.98959, 0.98958, 0.98838, 0.98835, 0.98826, 0.98825, 0.9883, 0.98828, 0.98711, 0.98709, 0.98562, 0.98561, 0.98511, 0.9851, 0.98414, 0.98412, 0.98284, 0.98282, 0.98104, 0.98101, 0.97981, 0.97979, 0.97849, 0.97849, 0.97557, 0.97554, 0.97398, 0.97395, 0.97171, 0.97166, 0.96917, 0.96913, 0.96511, 0.96507, 0.96263, 0.96257, 0.95839, 0.95835, 0.95483, 0.95475, 0.94942, 0.94936, 0.9468, 0.94678, 0.94583, 0.94594, 0.94843, 0.94872, 0.96949, 0.97015]),
|
||||
"i2v_14B_2.2": np.array([1.0]*2+[0.99512, 0.99559, 0.99559, 0.99561, 0.99595, 0.99577, 0.99512, 0.99512, 0.99546, 0.99534, 0.99543, 0.99531, 0.99496, 0.99491, 0.99504, 0.99499, 0.99444, 0.99449, 0.99481, 0.99481, 0.99435, 0.99435, 0.9943, 0.99431, 0.99411, 0.99406, 0.99373, 0.99376, 0.99413, 0.99405, 0.99363, 0.99359, 0.99335, 0.99331, 0.99244, 0.99243, 0.99229, 0.99229, 0.99239, 0.99236, 0.99163, 0.9916, 0.99149, 0.99151, 0.99191, 0.99192, 0.9898, 0.98981, 0.9899, 0.98987, 0.98849, 0.98849, 0.98846, 0.98846, 0.98861, 0.98861, 0.9874, 0.98738, 0.98588, 0.98589, 0.98539, 0.98534, 0.98444, 0.98439, 0.9831, 0.98309, 0.98119, 0.98118, 0.98001, 0.98, 0.97862, 0.97859, 0.97555, 0.97558, 0.97392, 0.97388, 0.97152, 0.97145, 0.96871, 0.9687, 0.96435, 0.96434, 0.96129, 0.96127, 0.95639, 0.95638, 0.95176, 0.95175, 0.94446, 0.94452, 0.93972, 0.93974, 0.93575, 0.9359, 0.93537, 0.93552, 0.96655, 0.96616]),
|
||||
}
|
||||
|
||||
model_variant = "14B" #default to this
|
||||
@@ -1242,6 +1255,13 @@ class WanVideoModelLoader:
|
||||
model_variant = "1_3B"
|
||||
if dim == 3072:
|
||||
log.info(f"5B model detected, no Teacache or MagCache coefficients available, consider using EasyCache for this model")
|
||||
|
||||
if "high" in model.lower() or "low" in model.lower():
|
||||
if "i2v" in model.lower():
|
||||
model_variant = "i2v_14B_2.2"
|
||||
else:
|
||||
model_variant = "14B_2.2"
|
||||
|
||||
log.info(f"Model variant detected: {model_variant}")
|
||||
|
||||
TRANSFORMER_CONFIG= {
|
||||
|
||||
+13
-2
@@ -12,7 +12,7 @@ from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_ti
|
||||
from .gguf.gguf import set_lora_params_gguf
|
||||
from .multitalk.multitalk import timestep_transform, add_noise
|
||||
from .utils import(log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, optimized_scale, setup_radial_attention,
|
||||
compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device, get_raag_guidance)
|
||||
compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device, get_raag_guidance, temporal_score_rescaling)
|
||||
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
|
||||
@@ -936,7 +936,7 @@ class WanVideoSampler:
|
||||
timesteps[-drift_steps:] = drift_timesteps[-drift_steps:]
|
||||
|
||||
# Experimental args
|
||||
use_cfg_zero_star = use_tangential = use_fresca = bidirectional_sampling =False
|
||||
use_cfg_zero_star = use_tangential = use_fresca = bidirectional_sampling = use_tsr = False
|
||||
raag_alpha = 0.0
|
||||
if experimental_args is not None:
|
||||
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
|
||||
@@ -960,6 +960,9 @@ class WanVideoSampler:
|
||||
bidirectional_sampling = experimental_args.get("bidirectional_sampling", False)
|
||||
if bidirectional_sampling:
|
||||
sample_scheduler_flipped = copy.deepcopy(sample_scheduler)
|
||||
use_tsr = experimental_args.get("temporal_score_rescaling", False)
|
||||
tsr_k = experimental_args.get("tsr_k", 1.0)
|
||||
tsr_sigma = experimental_args.get("tsr_sigma", 1.0)
|
||||
|
||||
# Rotary positional embeddings (RoPE)
|
||||
|
||||
@@ -2223,6 +2226,8 @@ class WanVideoSampler:
|
||||
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
|
||||
@@ -2693,6 +2698,9 @@ class WanVideoSampler:
|
||||
sampling_pbar.update(1)
|
||||
step_iteration_count += 1
|
||||
|
||||
if use_tsr:
|
||||
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
|
||||
|
||||
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
|
||||
|
||||
@@ -2795,6 +2803,9 @@ class WanVideoSampler:
|
||||
if flowedit_args is None:
|
||||
latent = latent.to(intermediate_device)
|
||||
|
||||
if use_tsr:
|
||||
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
|
||||
|
||||
if len(timestep.shape) != 1 and not is_pusa: #5b
|
||||
# all_indices is a list of indices to skip
|
||||
total_indices = list(range(latent.shape[1]))
|
||||
|
||||
@@ -68,12 +68,21 @@ EchoShot: https://github.com/D2I-ai/EchoShot
|
||||
|
||||
Stand-In: https://github.com/WeChatCV/Stand-In
|
||||
|
||||
HuMo: https://github.com/Phantom-video/HuMo
|
||||
|
||||
WanAnimate: https://github.com/Wan-Video/Wan2.2/tree/main/wan/modules/animate
|
||||
|
||||
|
||||
Examples:
|
||||
---
|
||||
|
||||
WanAnimate:
|
||||
|
||||
https://github.com/user-attachments/assets/f370b001-0f98-4c4c-bcb5-cfad0b330697
|
||||
|
||||
[ReCamMaster](https://github.com/KwaiVGI/ReCamMaster):
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/c58a12c2-13ba-4af8-8041-e283dbef197e
|
||||
|
||||
|
||||
|
||||
@@ -588,3 +588,16 @@ def check_duplicate_nodes():
|
||||
wanvideo_dirs.append(str(path))
|
||||
|
||||
return wanvideo_dirs
|
||||
|
||||
#https://github.com/temporalscorerescaling/TSR/
|
||||
def temporal_score_rescaling(model_output, sample, timestep, k=1.0, tsr_sigma=0.1):
|
||||
t = (timestep / 1000)
|
||||
if t == 0.0:
|
||||
ratio = k
|
||||
else:
|
||||
snr_t = (1 - t)**2 / t**2
|
||||
ratio = (snr_t * tsr_sigma**2 + 1) / (snr_t * tsr_sigma**2 / k + 1)
|
||||
|
||||
if not t == 1.0:
|
||||
model_output = (ratio * ((1-t) * model_output + sample) - sample) / (1 - t)
|
||||
return model_output
|
||||
|
||||
Reference in New Issue
Block a user