diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index 6284e64..d8e5b2b 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -34,7 +34,7 @@ from diffusers.schedulers import DPMSolverMultistepScheduler from ...modules import HYVideoDiffusionTransformer from comfy.utils import ProgressBar import math -from ....utils import optimized_scale +from ....utils import optimized_scale, fourier_filter logger = logging.get_logger(__name__) # pylint: disable=invalid-name EXAMPLE_DOC_STRING = """""" @@ -428,6 +428,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): sigmas: List[float] = None, guidance_scale: float = 1.0, use_cfg_zero_star: bool = False, + fresca_args: Optional[Dict[str, Any]] = None, cfg_start_percent: float = 0.0, cfg_end_percent: float = 1.0, batched_cfg: bool = True, @@ -727,6 +728,11 @@ class HunyuanVideoPipeline(DiffusionPipeline): else mask_latents ) 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") @@ -964,9 +970,18 @@ class HunyuanVideoPipeline(DiffusionPipeline): ).view(batch_size, 1, 1, 1) else: alpha = 1.0 - noise_pred = uncond * alpha + self.guidance_scale * (cond - uncond * alpha) + #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) - elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance: raise NotImplementedError 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_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: #reverse latent shift if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent: diff --git a/nodes.py b/nodes.py index cc2f08a..8cc485c 100644 --- a/nodes.py +++ b/nodes.py @@ -1233,6 +1233,25 @@ class HyVideoLoopArgs: def process(self, **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 class HyVideoSampler: @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"}), "i2v_mode": (["stability", "dynamic"], {"default": "dynamic", "tooltip": "I2V mode for image2video process"}), "loop_args": ("LOOPARGS", ), + "fresca_args": ("FRESCA_ARGS", ), "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, 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 device = mm.get_torch_device() @@ -1445,6 +1465,7 @@ class HyVideoSampler: cfg_end_percent=cfg_end_percent, batched_cfg=batched_cfg, use_cfg_zero_star=use_cfg_zero_star, + fresca_args=fresca_args, embedded_guidance_scale=embedded_guidance_scale, latents=input_latents, mask_latents=mask_latents, @@ -1884,6 +1905,7 @@ NODE_CLASS_MAPPINGS = { "HyVideoEncodeKeyframes": HyVideoEncodeKeyframes, "HyVideoTextEmbedBridge": HyVideoTextEmbedBridge, "HyVideoLoopArgs": HyVideoLoopArgs, + "HunyuanVideoFresca": HunyuanVideoFresca } NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoSampler": "HunyuanVideo Sampler", @@ -1912,4 +1934,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoEncodeKeyframes": "HyVideo Encode Keyframes", "HyVideoTextEmbedBridge": "HyVideo TextEmbed Bridge", "HyVideoLoopArgs": "HyVideo Loop Args", + "HunyuanVideoFresca": "HunyuanVideo Fresca" } diff --git a/utils.py b/utils.py index 4e0ae46..39a4c15 100644 --- a/utils.py +++ b/utils.py @@ -36,4 +36,57 @@ def optimized_scale(positive_flat, negative_flat): # st_star = v_cond^T * v_uncond / ||v_uncond||^2 st_star = dot_product / squared_norm - return st_star \ No newline at end of file + 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 \ No newline at end of file