Add FreSca

#https://github.com/WikiChao/FreSca
This commit is contained in:
kijai
2025-05-09 13:44:10 +03:00
parent 02ff79b3f4
commit 00c8864900
3 changed files with 104 additions and 5 deletions
@@ -34,7 +34,7 @@ from diffusers.schedulers import DPMSolverMultistepScheduler
from ...modules import HYVideoDiffusionTransformer from ...modules import HYVideoDiffusionTransformer
from comfy.utils import ProgressBar from comfy.utils import ProgressBar
import math import math
from ....utils import optimized_scale from ....utils import optimized_scale, fourier_filter
logger = logging.get_logger(__name__) # pylint: disable=invalid-name logger = logging.get_logger(__name__) # pylint: disable=invalid-name
EXAMPLE_DOC_STRING = """""" EXAMPLE_DOC_STRING = """"""
@@ -428,6 +428,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
sigmas: List[float] = None, sigmas: List[float] = None,
guidance_scale: float = 1.0, guidance_scale: float = 1.0,
use_cfg_zero_star: bool = False, use_cfg_zero_star: bool = False,
fresca_args: Optional[Dict[str, Any]] = None,
cfg_start_percent: float = 0.0, cfg_start_percent: float = 0.0,
cfg_end_percent: float = 1.0, cfg_end_percent: float = 1.0,
batched_cfg: bool = True, batched_cfg: bool = True,
@@ -728,6 +729,11 @@ class HunyuanVideoPipeline(DiffusionPipeline):
) )
print(f'mask_latents_model_input={mask_latents_model_input.shape} ') print(f'mask_latents_model_input={mask_latents_model_input.shape} ')
if fresca_args is not None:
fresca_scale_low = fresca_args.get("fresca_scale_low", 1.0)
fresca_scale_high = fresca_args.get("fresca_scale_high", 1.25)
fresca_freq_cutoff = fresca_args.get("fresca_freq_cutoff", 20)
logger.info(f"Sampling {video_length} frames in {latents.shape[2]} latents at {width}x{height} with {len(timesteps)} inference steps") logger.info(f"Sampling {video_length} frames in {latents.shape[2]} latents at {width}x{height} with {len(timesteps)} inference steps")
comfy_pbar = ProgressBar(len(timesteps)) comfy_pbar = ProgressBar(len(timesteps))
@@ -964,9 +970,18 @@ class HunyuanVideoPipeline(DiffusionPipeline):
).view(batch_size, 1, 1, 1) ).view(batch_size, 1, 1, 1)
else: else:
alpha = 1.0 alpha = 1.0
#https://github.com/WikiChao/FreSca
if fresca_args is not None:
filtered_cond = fourier_filter(
cond - uncond,
scale_low=fresca_scale_low,
scale_high=fresca_scale_high,
freq_cutoff=fresca_freq_cutoff,
)
noise_pred = uncond * alpha + self.guidance_scale * filtered_cond * alpha
else:
noise_pred = uncond * alpha + self.guidance_scale * (cond - uncond * alpha) noise_pred = uncond * alpha + self.guidance_scale * (cond - uncond * alpha)
elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance: elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance:
raise NotImplementedError raise NotImplementedError
noise_pred_uncond, noise_pred_text, noise_pred_perturb = noise_pred.chunk(3) noise_pred_uncond, noise_pred_text, noise_pred_perturb = noise_pred.chunk(3)
@@ -980,6 +995,14 @@ class HunyuanVideoPipeline(DiffusionPipeline):
noise_pred = noise_pred_text + self._stg_scale * ( noise_pred = noise_pred_text + self._stg_scale * (
noise_pred_text - noise_pred_perturb noise_pred_text - noise_pred_perturb
) )
else:
if fresca_args is not None:
noise_pred = fourier_filter(
noise_pred,
scale_low=fresca_scale_low,
scale_high=fresca_scale_high,
freq_cutoff=fresca_freq_cutoff,
)
if latent_shift_loop: if latent_shift_loop:
#reverse latent shift #reverse latent shift
if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent: if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent:
+24 -1
View File
@@ -1233,6 +1233,25 @@ class HyVideoLoopArgs:
def process(self, **kwargs): def process(self, **kwargs):
return (kwargs,) return (kwargs,)
class HunyuanVideoFresca:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"fresca_scale_low": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
"fresca_scale_high": ("FLOAT", {"default": 1.25, "min": 0.0, "max": 10.0, "step": 0.01}),
"fresca_freq_cutoff": ("INT", {"default": 20, "min": 0, "max": 10000, "step": 1}),
},
}
RETURN_TYPES = ("FRESCA_ARGS", )
RETURN_NAMES = ("fresca_args",)
FUNCTION = "process"
CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "https://github.com/WikiChao/FreSca"
def process(self, **kwargs):
return (kwargs,)
#region Sampler #region Sampler
class HyVideoSampler: class HyVideoSampler:
@classmethod @classmethod
@@ -1267,6 +1286,7 @@ class HyVideoSampler:
"riflex_freq_index": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Frequency index for RIFLEX, disabled when 0, default 4. Allows for new frames to be generated after 129 without looping"}), "riflex_freq_index": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Frequency index for RIFLEX, disabled when 0, default 4. Allows for new frames to be generated after 129 without looping"}),
"i2v_mode": (["stability", "dynamic"], {"default": "dynamic", "tooltip": "I2V mode for image2video process"}), "i2v_mode": (["stability", "dynamic"], {"default": "dynamic", "tooltip": "I2V mode for image2video process"}),
"loop_args": ("LOOPARGS", ), "loop_args": ("LOOPARGS", ),
"fresca_args": ("FRESCA_ARGS", ),
"mask": ("MASK", ), "mask": ("MASK", ),
} }
} }
@@ -1278,7 +1298,7 @@ class HyVideoSampler:
def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames, def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames,
samples=None, denoise_strength=1.0, force_offload=True, stg_args=None, context_options=None, feta_args=None, samples=None, denoise_strength=1.0, force_offload=True, stg_args=None, context_options=None, feta_args=None,
teacache_args=None, scheduler=None, image_cond_latents=None, neg_image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability", loop_args=None, mask=None): teacache_args=None, scheduler=None, image_cond_latents=None, neg_image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability", loop_args=None, fresca_args=None, mask=None):
model = model.model model = model.model
device = mm.get_torch_device() device = mm.get_torch_device()
@@ -1445,6 +1465,7 @@ class HyVideoSampler:
cfg_end_percent=cfg_end_percent, cfg_end_percent=cfg_end_percent,
batched_cfg=batched_cfg, batched_cfg=batched_cfg,
use_cfg_zero_star=use_cfg_zero_star, use_cfg_zero_star=use_cfg_zero_star,
fresca_args=fresca_args,
embedded_guidance_scale=embedded_guidance_scale, embedded_guidance_scale=embedded_guidance_scale,
latents=input_latents, latents=input_latents,
mask_latents=mask_latents, mask_latents=mask_latents,
@@ -1884,6 +1905,7 @@ NODE_CLASS_MAPPINGS = {
"HyVideoEncodeKeyframes": HyVideoEncodeKeyframes, "HyVideoEncodeKeyframes": HyVideoEncodeKeyframes,
"HyVideoTextEmbedBridge": HyVideoTextEmbedBridge, "HyVideoTextEmbedBridge": HyVideoTextEmbedBridge,
"HyVideoLoopArgs": HyVideoLoopArgs, "HyVideoLoopArgs": HyVideoLoopArgs,
"HunyuanVideoFresca": HunyuanVideoFresca
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoSampler": "HunyuanVideo Sampler", "HyVideoSampler": "HunyuanVideo Sampler",
@@ -1912,4 +1934,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoEncodeKeyframes": "HyVideo Encode Keyframes", "HyVideoEncodeKeyframes": "HyVideo Encode Keyframes",
"HyVideoTextEmbedBridge": "HyVideo TextEmbed Bridge", "HyVideoTextEmbedBridge": "HyVideo TextEmbed Bridge",
"HyVideoLoopArgs": "HyVideo Loop Args", "HyVideoLoopArgs": "HyVideo Loop Args",
"HunyuanVideoFresca": "HunyuanVideo Fresca"
} }
+53
View File
@@ -37,3 +37,56 @@ def optimized_scale(positive_flat, negative_flat):
st_star = dot_product / squared_norm st_star = dot_product / squared_norm
return st_star return st_star
# Code based on https://github.com/WikiChao/FreSca (MIT License)
import torch
import torch.fft as fft
def fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
"""
Apply frequency-dependent scaling to an image tensor using Fourier transforms.
Parameters:
x: Input tensor of shape (B, C, H, W)
scale_low: Scaling factor for low-frequency components (default: 1.0)
scale_high: Scaling factor for high-frequency components (default: 1.5)
freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20)
Returns:
x_filtered: Filtered version of x in spatial domain with frequency-specific scaling applied.
"""
# Preserve input dtype and device
dtype, device = x.dtype, x.device
# Convert to float32 for FFT computations
x = x.to(torch.float32)
# 1) Apply FFT and shift low frequencies to center
x_freq = fft.fftn(x, dim=(-2, -1))
x_freq = fft.fftshift(x_freq, dim=(-2, -1))
# 2) Create a mask to scale frequencies differently
C, B, H, W = x_freq.shape
crow, ccol = H // 2, W // 2
# Initialize mask with high-frequency scaling factor
mask = torch.ones((C, B, H, W), device=device) * scale_high
# Apply low-frequency scaling factor to center region
mask[
...,
crow - freq_cutoff : crow + freq_cutoff,
ccol - freq_cutoff : ccol + freq_cutoff,
] = scale_low
# 3) Apply frequency-specific scaling
x_freq = x_freq * mask
# 4) Convert back to spatial domain
x_freq = fft.ifftshift(x_freq, dim=(-2, -1))
x_filtered = fft.ifftn(x_freq, dim=(-2, -1)).real
# 5) Restore original dtype
x_filtered = x_filtered.to(dtype)
return x_filtered