@@ -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:
|
||||||
|
|||||||
@@ -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"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
Reference in New Issue
Block a user