FreeInit
https: //github.com/TianxingWu/FreeInit Co-Authored-By: kabachuha <14872007+kabachuha@users.noreply.github.com>
This commit is contained in:
@@ -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
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user