@@ -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,
|
||||
@@ -728,6 +729,11 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
)
|
||||
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")
|
||||
|
||||
comfy_pbar = ProgressBar(len(timesteps))
|
||||
@@ -964,9 +970,18 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
).view(batch_size, 1, 1, 1)
|
||||
else:
|
||||
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)
|
||||
|
||||
|
||||
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:
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -37,3 +37,56 @@ def optimized_scale(positive_flat, negative_flat):
|
||||
st_star = dot_product / squared_norm
|
||||
|
||||
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
|
||||
Reference in New Issue
Block a user