https: //github.com/TianxingWu/FreeInit
Co-Authored-By: kabachuha <14872007+kabachuha@users.noreply.github.com>
This commit is contained in:
kijai
2025-07-08 15:09:54 +03:00
co-authored by kabachuha
parent 8500514ef9
commit da98636599
2 changed files with 729 additions and 487 deletions
+142
View File
@@ -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
+102 -2
View File
@@ -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
@@ -2063,6 +2086,24 @@ class WanVideoSampler:
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"])
dwpose_data = unianimate_poses["pose"].to(device, model["dtype"])
@@ -2724,6 +2765,52 @@ class WanVideoSampler:
except:
pass
# 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
print(latent)
#region main loop start
for idx, t in enumerate(tqdm(timesteps)):
if flowedit_args is not None:
@@ -3268,6 +3355,12 @@ class WanVideoSampler:
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)
@@ -3317,7 +3410,12 @@ 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:
@@ -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",
}