18 Commits
Author SHA1 Message Date
kijai 3ce9640497 Update nodes.py 2025-05-09 20:13:08 +03:00
kijai 675baf7574 fix 2025-05-09 16:56:29 +03:00
kijai 9a5d192807 more TeaCache fixes.. 2025-05-09 16:43:35 +03:00
kijai 14cc8ae0a5 TeaCache fixes 2025-05-09 16:36:00 +03:00
kijai 0d9fe9b86a Update nodes.py 2025-05-09 16:34:50 +03:00
kijai 27e541b0a5 Update models.py 2025-05-09 15:05:40 +03:00
kijai 9b321d09c5 Update nodes.py 2025-05-09 15:02:00 +03:00
kijai 91aabfc857 fixes 2025-05-09 14:59:44 +03:00
kijai 6682a0905e Add SLG
For testing, no clue what blocks to use at this point
2025-05-09 14:10:55 +03:00
kijai f5a243a898 Update utils.py 2025-05-09 13:48:30 +03:00
kijai 00c8864900 Add FreSca
#https://github.com/WikiChao/FreSca
2025-05-09 13:44:10 +03:00
kijai 02ff79b3f4 Update hyvideo_custom_testing_01.json 2025-05-09 13:01:38 +03:00
kijai e18afda414 Add TeaCache start/end step 2025-05-09 11:41:31 +03:00
kijai 64b5d31765 Update hyvideo_custom_testing_01.json 2025-05-09 09:45:22 +03:00
kijai 468b103d38 Create hyvideo_custom_testing_01.json 2025-05-09 09:37:42 +03:00
kijai 6ea4d31b41 initial hunyuan custom support 2025-05-09 09:16:02 +03:00
kijai 83f0bbb869 Update nodes.py 2025-03-30 19:48:07 +03:00
kijai 75190b756a Add Mobius looping option
https://github.com/YisuiTT/Mobius/
2025-03-30 19:07:04 +03:00
23 changed files with 1939 additions and 1385 deletions
File diff suppressed because it is too large Load Diff
+3 -2
View File
@@ -7,8 +7,9 @@ def fp8_linear_forward(cls, original_dtype, input):
weight_dtype = cls.weight.dtype weight_dtype = cls.weight.dtype
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
if len(input.shape) == 3: if len(input.shape) == 3:
target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn #target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn
inn = input.reshape(-1, input.shape[2]).to(target_dtype) #inn = input.reshape(-1, input.shape[2]).to(target_dtype)
inn = input.reshape(-1, input.shape[2]).to(weight_dtype)
w = cls.weight.t() w = cls.weight.t()
scale = torch.ones((1), device=input.device, dtype=torch.float32) scale = torch.ones((1), device=input.device, dtype=torch.float32)
@@ -33,14 +33,15 @@ 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
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 = """"""
from ...modules.posemb_layers import get_nd_rotary_pos_embed from ...modules.posemb_layers import get_nd_rotary_pos_embed, get_nd_rotary_pos_embed_new
from ....enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight from ....enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight
def get_rotary_pos_embed(transformer, latent_video_length, height, width, k=0): def get_rotary_pos_embed(transformer, latent_video_length, height, width, k=0, rope_func=get_nd_rotary_pos_embed):
target_ndim = 3 target_ndim = 3
ndim = 5 - 2 ndim = 5 - 2
rope_theta = 225 rope_theta = 225
@@ -79,7 +80,7 @@ def get_rotary_pos_embed(transformer, latent_video_length, height, width, k=0):
assert ( assert (
sum(rope_dim_list) == head_dim sum(rope_dim_list) == head_dim
), "sum(rope_dim_list) should equal to head_dim of attention layer" ), "sum(rope_dim_list) should equal to head_dim of attention layer"
freqs_cos, freqs_sin = get_nd_rotary_pos_embed( freqs_cos, freqs_sin = rope_func(
rope_dim_list, rope_dim_list,
rope_sizes, rope_sizes,
theta=rope_theta, theta=rope_theta,
@@ -254,6 +255,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
) )
if latents is not None: if latents is not None:
latents = latents.to(device) latents = latents.to(device)
else:
original_latents = None
noise = randn_tensor(shape, generator=generator, device=device, dtype=self.base_dtype) noise = randn_tensor(shape, generator=generator, device=device, dtype=self.base_dtype)
if freenoise: if freenoise:
@@ -319,6 +322,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latents = latents[:, :, :frames_needed, :, :] latents = latents[:, :, :frames_needed, :, :]
logger.info(f"Frames needed less than current frames, cutting down to {frames_needed}") logger.info(f"Frames needed less than current frames, cutting down to {frames_needed}")
original_latents = latents.clone()
latents = latents * (1 - latent_timestep / 1000) + latent_timestep / 1000 * noise latents = latents * (1 - latent_timestep / 1000) + latent_timestep / 1000 * noise
print("latents shape:", latents.shape) print("latents shape:", latents.shape)
@@ -338,7 +342,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if hasattr(self.scheduler, "init_noise_sigma"): if hasattr(self.scheduler, "init_noise_sigma"):
# scale the initial noise by the standard deviation required by the scheduler # scale the initial noise by the standard deviation required by the scheduler
latents = latents * self.scheduler.init_noise_sigma latents = latents * self.scheduler.init_noise_sigma
return latents.to(device), timesteps, i2v_mask, image_cond_latents return latents.to(device), timesteps, i2v_mask, image_cond_latents, noise, original_latents
# Copied from diffusers.pipelines.latent_consistency_models.pipeline_latent_consistency_text2img.LatentConsistencyModelPipeline.get_guidance_scale_embedding # Copied from diffusers.pipelines.latent_consistency_models.pipeline_latent_consistency_text2img.LatentConsistencyModelPipeline.get_guidance_scale_embedding
def get_guidance_scale_embedding( def get_guidance_scale_embedding(
@@ -423,6 +427,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
timesteps: List[int] = None, timesteps: List[int] = None,
sigmas: List[float] = None, sigmas: List[float] = None,
guidance_scale: float = 1.0, guidance_scale: float = 1.0,
use_cfg_zero_star: bool = False,
fresca_args: Optional[Dict[str, Any]] = None,
slg_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,
@@ -431,6 +438,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
denoise_strength: float = 1.0, denoise_strength: float = 1.0,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None, latents: Optional[torch.Tensor] = None,
mask_latents: Optional[torch.Tensor] = None,
cross_attention_kwargs: Optional[Dict[str, Any]] = None, cross_attention_kwargs: Optional[Dict[str, Any]] = None,
guidance_rescale: float = 0.0, guidance_rescale: float = 0.0,
clip_skip: Optional[int] = None, clip_skip: Optional[int] = None,
@@ -452,9 +460,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
feta_args: Optional[Dict] = None, feta_args: Optional[Dict] = None,
leapfusion_img2vid: Optional[bool] = False, leapfusion_img2vid: Optional[bool] = False,
image_cond_latents: Optional[torch.Tensor] = None, image_cond_latents: Optional[torch.Tensor] = None,
neg_image_cond_latents: Optional[torch.Tensor] = None,
riflex_freq_index: Optional[int] = None, riflex_freq_index: Optional[int] = None,
i2v_stability=True, i2v_stability=True,
taylorseer: Optional[dict] = None, loop_args: Optional[Dict] = None,
**kwargs, **kwargs,
): ):
r""" r"""
@@ -540,6 +549,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
# 2. Define call parameters # 2. Define call parameters
batch_size = 1 batch_size = 1
ref_latents = None
device = self._execution_device device = self._execution_device
prompt_embeds = prompt_embed_dict.get("prompt_embeds", None) prompt_embeds = prompt_embed_dict.get("prompt_embeds", None)
@@ -634,14 +644,25 @@ class HunyuanVideoPipeline(DiffusionPipeline):
use_context_schedule = True use_context_schedule = True
from ....context import get_context_scheduler from ....context import get_context_scheduler
context = get_context_scheduler(context_schedule) context = get_context_scheduler(context_schedule)
freqs_cos, freqs_sin = get_rotary_pos_embed( if i2v_condition_type == "reference":
self.transformer, context_frames, height, width freqs_cos, freqs_sin = get_rotary_pos_embed(
) self.transformer, context_frames, height, width, rope_func=get_nd_rotary_pos_embed_new
)
else:
freqs_cos, freqs_sin = get_rotary_pos_embed(
self.transformer, context_frames, height, width
)
else: else:
# rotary embeddings # rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed( if i2v_condition_type == "reference":
self.transformer, latent_video_length, height, width, k=riflex_freq_index print("Using reference condition")
) freqs_cos, freqs_sin = get_rotary_pos_embed(
self.transformer, latent_video_length, height, width, rope_func=get_nd_rotary_pos_embed_new
)
else:
freqs_cos, freqs_sin = get_rotary_pos_embed(
self.transformer, latent_video_length, height, width, k=riflex_freq_index
)
if not self.transformer.upcast_rope: if not self.transformer.upcast_rope:
freqs_cos = freqs_cos.to(self.base_dtype).to(device) freqs_cos = freqs_cos.to(self.base_dtype).to(device)
freqs_sin = freqs_sin.to(self.base_dtype).to(device) freqs_sin = freqs_sin.to(self.base_dtype).to(device)
@@ -652,10 +673,11 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if leapfusion_img2vid: if leapfusion_img2vid:
logger.info("Single input latent frame detected, LeapFusion img2vid enabled") logger.info("Single input latent frame detected, LeapFusion img2vid enabled")
original_latents = latents original_latents = latents
# 5. Prepare latent variables # 5. Prepare latent variables
#num_channels_latents = self.transformer.config.in_channels #num_channels_latents = self.transformer.config.in_channels
num_channels_latents = 16 num_channels_latents = 16
latents, timesteps, i2v_mask, image_cond_latents = self.prepare_latents( latents, timesteps, i2v_mask, image_cond_latents, noise, original_latents = self.prepare_latents(
batch_size * num_videos_per_prompt, batch_size * num_videos_per_prompt,
num_channels_latents, num_channels_latents,
num_inference_steps, num_inference_steps,
@@ -691,13 +713,36 @@ class HunyuanVideoPipeline(DiffusionPipeline):
#print(self.scheduler.sigmas) #print(self.scheduler.sigmas)
tseercache_dict, tseer_current = None, None latent_shift_loop = False
if taylorseer: if loop_args is not None:
print(taylorseer) latent_shift_loop = True
from ...modules.cache_functions import cache_init is_looped = True
tseercache_dict, tseer_current = cache_init(self._num_timesteps, cache_device=taylorseer["cache_device"], compute_device=taylorseer["compute_device"]) latent_skip = loop_args["shift_skip"]
tseercache_dict["max_order"] = taylorseer["max_order"] latent_shift_start_percent = loop_args["start_percent"]
tseercache_dict["fresh_threshold"] = taylorseer["fresh_threshold"] latent_shift_end_percent = loop_args["end_percent"]
shift_idx = 0
if mask_latents is not None:
mask_latents_model_input = (
torch.cat([mask_latents] * 2)
if not math.isclose(self.guidance_scale, 1.0)
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)
if slg_args is not None:
assert batched_cfg is not None, "Batched cfg is not supported with SLG"
self.transformer.slg_single_blocks = slg_args["single_blocks"]
self.transformer.slg_double_blocks = slg_args["double_blocks"]
self.transformer.slg_start_percent = slg_args["start_percent"]
self.transformer.slg_end_percent = slg_args["end_percent"]
else:
self.transformer.slg_single_blocks = self.transformer.slg_double_blocks = None
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")
@@ -707,8 +752,16 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if self.interrupt: if self.interrupt:
continue continue
current_step_percentage = i / len(timesteps)
if image_cond_latents is not None and i2v_condition_type == "token_replace": if image_cond_latents is not None and i2v_condition_type == "token_replace":
latents = torch.concat([original_image_latents, latents[:, :, 1:, :, :]], dim=2) latents = torch.concat([original_image_latents, latents[:, :, 1:, :, :]], dim=2)
elif image_cond_latents is not None and i2v_condition_type == "reference":
ref_latents = image_cond_latents
if neg_image_cond_latents is not None:
uncond_ref_latents = neg_image_cond_latents
else:
uncond_ref_latents = image_cond_latents
latent_model_input = latents latent_model_input = latents
input_prompt_embeds = prompt_embeds input_prompt_embeds = prompt_embeds
@@ -717,10 +770,11 @@ class HunyuanVideoPipeline(DiffusionPipeline):
cfg_enabled = False cfg_enabled = False
stg_enabled = False stg_enabled = False
current_step_percentage = i / len(timesteps) ### latent shift
if latent_shift_loop:
if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent:
latent_model_input = torch.cat([latent_model_input[:, :, shift_idx:]] + [latent_model_input[:, :, :shift_idx]], dim=2)
if taylorseer:
tseer_current['step'] = i
if self.do_spatio_temporal_guidance: if self.do_spatio_temporal_guidance:
if stg_start_percent <= current_step_percentage <= stg_end_percent: if stg_start_percent <= current_step_percentage <= stg_end_percent:
@@ -758,6 +812,16 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
if mask_latents is not None:
original_latents_noise = original_latents * (1 - t / 1000.0) + t / 1000.0 * noise
original_latent_noise_model_input = (
torch.cat([original_latents_noise] * 2)
if self.do_classifier_free_guidance
else original_latents_noise
)
original_latent_noise_model_input = self.scheduler.scale_model_input(original_latent_noise_model_input, t)
latent_model_input = mask_latents_model_input * latent_model_input + (1 - mask_latents_model_input) * original_latent_noise_model_input
t_expand = t.repeat(latent_model_input.shape[0]) t_expand = t.repeat(latent_model_input.shape[0])
if leapfusion_img2vid: if leapfusion_img2vid:
@@ -828,6 +892,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_block_idx=stg_block_idx, stg_block_idx=stg_block_idx,
stg_mode=stg_mode, stg_mode=stg_mode,
return_dict=True, return_dict=True,
ref_latents=ref_latents
)["x"] )["x"]
window_mask = torch.ones_like(noise_pred_context) window_mask = torch.ones_like(noise_pred_context)
@@ -864,8 +929,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_block_idx=stg_block_idx, stg_block_idx=stg_block_idx,
stg_mode=stg_mode, stg_mode=stg_mode,
return_dict=True, return_dict=True,
tseercache_dict = tseercache_dict, #taylorseer ref_latents=ref_latents,
tseer_current = tseer_current, #taylorseer is_uncond = False,
current_step = i,
current_step_percentage = current_step_percentage
)["x"] )["x"]
else: else:
uncond = self.transformer( uncond = self.transformer(
@@ -876,12 +943,14 @@ class HunyuanVideoPipeline(DiffusionPipeline):
text_states_2=input_prompt_embeds_2[0].unsqueeze(0), text_states_2=input_prompt_embeds_2[0].unsqueeze(0),
freqs_cos=freqs_cos, freqs_cos=freqs_cos,
freqs_sin=freqs_sin, freqs_sin=freqs_sin,
guidance=guidance_expand[0].unsqueeze(0), guidance=guidance_expand[0].unsqueeze(0) if guidance_expand is not None else None,
stg_block_idx=stg_block_idx, stg_block_idx=stg_block_idx,
stg_mode=stg_mode, stg_mode=stg_mode,
return_dict=True, return_dict=True,
tseercache_dict = tseercache_dict, #taylorseer ref_latents=uncond_ref_latents,
tseer_current = tseer_current, #taylorseer is_uncond = True,
current_step = i,
current_step_percentage = current_step_percentage
)["x"] )["x"]
cond = self.transformer( cond = self.transformer(
latent_model_input[1].unsqueeze(0), latent_model_input[1].unsqueeze(0),
@@ -891,24 +960,40 @@ class HunyuanVideoPipeline(DiffusionPipeline):
text_states_2=input_prompt_embeds_2[1].unsqueeze(0), text_states_2=input_prompt_embeds_2[1].unsqueeze(0),
freqs_cos=freqs_cos, freqs_cos=freqs_cos,
freqs_sin=freqs_sin, freqs_sin=freqs_sin,
guidance=guidance_expand[1].unsqueeze(0), guidance=guidance_expand[1].unsqueeze(0) if guidance_expand is not None else None,
stg_block_idx=stg_block_idx, stg_block_idx=stg_block_idx,
stg_mode=stg_mode, stg_mode=stg_mode,
return_dict=True, return_dict=True,
tseercache_dict = tseercache_dict, #taylorseer ref_latents=ref_latents,
tseer_current = tseer_current, #taylorseer is_uncond = False,
current_step = i,
current_step_percentage = current_step_percentage
)["x"] )["x"]
# perform guidance # perform guidance
if cfg_enabled and not self.do_spatio_temporal_guidance: if cfg_enabled and not self.do_spatio_temporal_guidance:
if batched_cfg: if batched_cfg:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) uncond, cond = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + self.guidance_scale * (
noise_pred_text - noise_pred_uncond
)
else:
noise_pred = uncond + self.guidance_scale * (cond - uncond)
#https://github.com/WeichenFan/CFG-Zero-star/
if use_cfg_zero_star:
alpha = optimized_scale(
cond.view(batch_size, -1),
uncond.view(batch_size, -1)
).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: elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance:
raise NotImplementedError raise NotImplementedError
@@ -923,6 +1008,19 @@ 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:
#reverse latent shift
if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent:
noise_pred = torch.cat([noise_pred[:, :, latent_video_length - shift_idx:]] + [noise_pred[:, :, :latent_video_length - shift_idx]], dim=2)
shift_idx = (shift_idx + latent_skip) % latent_video_length
# compute the previous noisy sample x_t -> x_t-1 # compute the previous noisy sample x_t -> x_t-1
if image_cond_latents is not None and i2v_condition_type == "token_replace": if image_cond_latents is not None and i2v_condition_type == "token_replace":
@@ -969,6 +1067,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
else: else:
comfy_pbar.update(1) comfy_pbar.update(1)
if mask_latents is not None:
latents = mask_latents * latents + (1 - mask_latents) * original_latents
if image_cond_latents is not None: if image_cond_latents is not None:
if leapfusion_img2vid or i2v_condition_type == "latent_concat": if leapfusion_img2vid or i2v_condition_type == "latent_concat":
latents = latents[:, :, 1:, :, :] latents = latents[:, :, 1:, :, :]
@@ -1,12 +0,0 @@
from .cache_cutfresh import cache_cutfresh
from .fresh_ratio_scheduler import fresh_ratio_scheduler
from .score_evaluate import score_evaluate
from .global_force_fresh import global_force_fresh
from .cache_cutfresh import cache_cutfresh
from .update_cache import update_cache
from .force_init import force_init
from .attention import cached_attention_forward
from .cache_init import cache_init
from .cal_type import cal_type
from .force_scheduler import force_scheduler
from .support_set_selection import support_set_selection
@@ -1,31 +0,0 @@
# Besides, re-arrange the attention module
from torch.jit import Final
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Union
#from xformers.ops.fmha.attn_bias import BlockDiagonalMask
def cached_attention_forward(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
#attn_bias: Optional[Union[torch.Tensor, BlockDiagonalMask]] = None,
attn_bias,
p: float = 0.0,
scale: Optional[float] = None
) -> torch.Tensor:
scale = 1.0 / query.shape[-1] ** 0.5
query = query * scale
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
attn = query @ key.transpose(-2, -1)
if attn_bias is not None:
attn_bias = attn_bias.materialize(shape= attn.shape, dtype= attn.dtype, device= attn.device)
attn = attn + attn_bias
#out_map = attn
attn_map = attn.softmax(-1)
attn = F.dropout(attn_map, p)
attn = attn @ value
return attn.transpose(1, 2).contiguous(), attn_map.mean(dim=1)
@@ -1,75 +0,0 @@
from .fresh_ratio_scheduler import fresh_ratio_scheduler
from .score_evaluate import score_evaluate
#from .token_merge import token_merge
from .support_set_selection import support_set_selection
import torch
def cache_cutfresh(cache_dic, tokens, current):
'''
Cut fresh tokens from the input tokens and update the cache counter.
cache_dic: dict, the cache dictionary containing cache(main extra memory cost), indices and some other information.
tokens: torch.Tensor, the input tokens to be cut.
current: dict, the current step, layer, and module information. Particularly convenient for debugging.
'''
step = current['step']
layer = current['layer']
stream = current['stream']
module = current['module']
fresh_ratio = fresh_ratio_scheduler(cache_dic, current)
fresh_ratio = torch.clamp(torch.tensor(fresh_ratio, device = tokens.device), min=0, max=1)
# Generate the index tensor for fresh tokens
score = score_evaluate(cache_dic, tokens, current) # s1, s2, s3 mentioned in the paper
#score = local_selection_with_bonus(score, 0.4, 4) # Uniform Spatial Distribution s4 mentioned in the paper
indices = score.argsort(dim=-1, descending=True)
topk = int(fresh_ratio * score.shape[1])
fresh_indices = indices[:, :topk]
stale_indices = indices[:, topk:]
#fresh_indices = support_set_selection(tokens, fresh_ratio, 0.4, current, cache_dic) # (B, fresh_ratio * N) # 0.4
# (B, fresh_ratio *N)
# Updating the Cache Frequency Score s3 mentioned in the paper
# stale tokens index + 1 in each ***module***, fresh tokens index = 0
###cache_dic['cache_index'][-1][layer][module] += 1
###cache_dic['cache_index'][-1][layer][module].scatter_(dim=1, index=fresh_indices,
### src = torch.zeros_like(fresh_indices, dtype=torch.int, device=fresh_indices.device))
#cache_dic['cache_index']['layer_index'][module] += 1
#cache_dic['cache_index']['layer_index'][module].scatter_(dim=1, index=fresh_indices,
# src = torch.zeros_like(fresh_indices, dtype=torch.int, device=fresh_indices.device))
fresh_indices_expand = fresh_indices.unsqueeze(-1).expand(-1, -1, tokens.shape[-1])
fresh_tokens = torch.gather(input = tokens, dim = 1, index = fresh_indices_expand)
return fresh_indices, fresh_tokens
def local_selection_with_bonus(score, bonus_ratio, grid_size=2):
batch_size, num_tokens = score.shape
image_size = int(num_tokens ** 0.5)
block_size = grid_size * grid_size
assert num_tokens % block_size == 0, "The number of tokens must be divisible by the block size."
# Step 1: Reshape score to group it by blocks
score_reshaped = score.view(batch_size, image_size // grid_size, grid_size, image_size // grid_size, grid_size)
score_reshaped = score_reshaped.permute(0, 1, 3, 2, 4).contiguous()
score_reshaped = score_reshaped.view(batch_size, -1, block_size) # [batch_size, num_blocks, block_size]
# Step 2: Find the max token in each block
max_scores, max_indices = score_reshaped.max(dim=-1, keepdim=True) # [batch_size, num_blocks, 1]
# Step 3: Create a mask to identify max score tokens
mask = torch.zeros_like(score_reshaped)
mask.scatter_(-1, max_indices, 1) # Set mask to 1 at the max indices
# Step 4: Apply the bonus only to the max score tokens
score_reshaped = score_reshaped + (mask * max_scores * bonus_ratio) # Apply bonus only to max tokens
# Step 5: Reshape the score back to its original shape
score_modified = score_reshaped.view(batch_size, image_size // grid_size, image_size // grid_size, grid_size, grid_size)
score_modified = score_modified.permute(0, 1, 3, 2, 4).contiguous()
score_modified = score_modified.view(batch_size, num_tokens)
return score_modified
@@ -1,124 +0,0 @@
import torch
def cache_init(num_steps, model_kwargs=None, cache_device=torch.device("cpu"), compute_device=torch.device("cuda")):
'''
Initialization for cache.
'''
cache_dic = {}
cache = {}
cache_index = {}
cache[-1]={}
cache_index[-1]={}
cache_index['layer_index']={}
cache_dic['attn_map'] = {}
cache_dic['attn_map'][-1] = {}
cache_dic['attn_map'][-1]['double_stream'] = {}
cache_dic['attn_map'][-1]['single_stream'] = {}
cache_dic['k-norm'] = {}
cache_dic['k-norm'][-1] = {}
cache_dic['k-norm'][-1]['double_stream'] = {}
cache_dic['k-norm'][-1]['single_stream'] = {}
cache_dic['v-norm'] = {}
cache_dic['v-norm'][-1] = {}
cache_dic['v-norm'][-1]['double_stream'] = {}
cache_dic['v-norm'][-1]['single_stream'] = {}
cache_dic['cross_attn_map'] = {}
cache_dic['cross_attn_map'][-1] = {}
cache[-1]['double_stream']={}
cache[-1]['single_stream']={}
cache_dic['cache_counter'] = 0
cache_dic['cache_device'] = cache_device
cache_dic['compute_device'] = compute_device
for j in range(20):
cache[-1]['double_stream'][j] = {}
cache_index[-1][j] = {}
cache_dic['attn_map'][-1]['double_stream'][j] = {}
cache_dic['attn_map'][-1]['double_stream'][j]['total'] = {}
cache_dic['attn_map'][-1]['double_stream'][j]['txt_mlp'] = {}
cache_dic['attn_map'][-1]['double_stream'][j]['img_mlp'] = {}
cache_dic['k-norm'][-1]['double_stream'][j] = {}
cache_dic['k-norm'][-1]['double_stream'][j]['txt_mlp'] = {}
cache_dic['k-norm'][-1]['double_stream'][j]['img_mlp'] = {}
cache_dic['v-norm'][-1]['double_stream'][j] = {}
cache_dic['v-norm'][-1]['double_stream'][j]['txt_mlp'] = {}
cache_dic['v-norm'][-1]['double_stream'][j]['img_mlp'] = {}
for j in range(40):
cache[-1]['single_stream'][j] = {}
cache_index[-1][j] = {}
cache_dic['attn_map'][-1]['single_stream'][j] = {}
cache_dic['attn_map'][-1]['single_stream'][j]['total'] = {}
cache_dic['k-norm'][-1]['single_stream'][j] = {}
cache_dic['k-norm'][-1]['single_stream'][j]['total'] = {}
cache_dic['v-norm'][-1]['single_stream'][j] = {}
cache_dic['v-norm'][-1]['single_stream'][j]['total'] = {}
cache_dic['taylor_cache'] = False
cache_dic['duca'] = False
cache_dic['test_FLOPs'] = False
mode = 'Taylor'
if mode == 'original':
cache_dic['cache_type'] = 'random'
cache_dic['cache_index'] = cache_index
cache_dic['cache'] = cache
cache_dic['fresh_ratio_schedule'] = 'ToCa'
cache_dic['fresh_ratio'] = 0.0
cache_dic['fresh_threshold'] = 1
cache_dic['force_fresh'] = 'global'
cache_dic['soft_fresh_weight'] = 0.0
cache_dic['max_order'] = 0
cache_dic['first_enhance'] = 1
elif mode == 'ToCa':
cache_dic['cache_type'] = 'random'
cache_dic['cache_index'] = cache_index
cache_dic['cache'] = cache
cache_dic['fresh_ratio_schedule'] = 'ToCa'
cache_dic['fresh_ratio'] = 0.10
cache_dic['fresh_threshold'] = 5
cache_dic['force_fresh'] = 'global'
cache_dic['soft_fresh_weight'] = 0.0
cache_dic['max_order'] = 0
cache_dic['first_enhance'] = 1
cache_dic['duca'] = False
elif mode == 'DuCa':
cache_dic['cache_type'] = 'random'
cache_dic['cache_index'] = cache_index
cache_dic['cache'] = cache
cache_dic['fresh_ratio_schedule'] = 'ToCa'
cache_dic['fresh_ratio'] = 0.10
cache_dic['fresh_threshold'] = 5
cache_dic['force_fresh'] = 'global'
cache_dic['soft_fresh_weight'] = 0.0
cache_dic['max_order'] = 0
cache_dic['first_enhance'] = 1
cache_dic['duca'] = True
elif mode == 'Taylor':
cache_dic['cache_type'] = 'random'
cache_dic['cache_index'] = cache_index
cache_dic['cache'] = cache
cache_dic['fresh_ratio_schedule'] = 'ToCa'
cache_dic['fresh_ratio'] = 0.0
cache_dic['fresh_threshold'] = 5
cache_dic['max_order'] = 1
cache_dic['force_fresh'] = 'global'
cache_dic['soft_fresh_weight'] = 0.0
cache_dic['taylor_cache'] = True
cache_dic['first_enhance'] = 1
current = {}
current['num_steps'] = num_steps
current['activated_steps'] = [0]
return cache_dic, current
@@ -1,49 +0,0 @@
from .force_scheduler import force_scheduler
def cal_type(cache_dic, current):
'''
Determine calculation type for this step
'''
if (cache_dic['fresh_ratio'] == 0.0) and (not cache_dic['taylor_cache']):
# FORA:Uniform
first_step = (current['step'] == 0)
else:
# ToCa: First enhanced
first_step = (current['step'] < cache_dic['first_enhance'])
#first_step = (current['step'] <= 3)
force_fresh = cache_dic['force_fresh']
if not first_step:
fresh_interval = cache_dic['cal_threshold']
else:
fresh_interval = cache_dic['fresh_threshold']
if (first_step) or (cache_dic['cache_counter'] == fresh_interval - 1 ):
current['type'] = 'full'
cache_dic['cache_counter'] = 0
current['activated_steps'].append(current['step'])
#current['activated_times'].append(current['t'])
force_scheduler(cache_dic, current)
elif (cache_dic['taylor_cache']):
cache_dic['cache_counter'] += 1
current['type'] = 'taylor_cache'
else:
cache_dic['cache_counter'] += 1
if (cache_dic['duca']):
if (cache_dic['cache_counter'] % 2 == 1): # 0: ToCa-Aggresive-ToCa, 1: Aggresive-ToCa-Aggresive
current['type'] = 'ToCa'
# 'cache_noise' 'ToCa' 'FORA'
else:
current['type'] = 'aggressive'
else:
current['type'] = 'ToCa'
#if current['step'] < 25:
# current['type'] = 'FORA'
#else:
# current['type'] = 'aggressive'
######################################################################
#if (current['step'] in [3,2,1,0]):
# current['type'] = 'full'
@@ -1,10 +0,0 @@
import torch
def force_init(cache_dic, current, tokens):
'''
Initialization for Force Activation step.
'''
cache_dic['cache_index'][-1][current['layer']][current['module']] = torch.zeros(tokens.shape[0], tokens.shape[1], dtype=torch.int, device=tokens.device)
#if current['layer'] == 0:
# cache_dic['cache_index']['layer_index'][current['module']] = torch.zeros(tokens.shape[0], tokens.shape[1], dtype=torch.int, device=tokens.device)
@@ -1,16 +0,0 @@
import torch
def force_scheduler(cache_dic, current):
if cache_dic['fresh_ratio'] == 0:
# FORA
linear_step_weight = 0.0
else:
# TokenCache
linear_step_weight = 0.0
step_factor = torch.tensor(1 - linear_step_weight + 2 * linear_step_weight * current['step'] / current['num_steps'])
threshold = torch.round(cache_dic['fresh_threshold'] / step_factor)
# no force constrain for sensitive steps, cause the performance is good enough.
# you may have a try.
cache_dic['cal_threshold'] = threshold
#return threshold
@@ -1,59 +0,0 @@
import torch
def fresh_ratio_scheduler(cache_dic, current):
'''
Return the fresh ratio for the current step.
'''
fresh_ratio = cache_dic['fresh_ratio']
fresh_ratio_schedule = cache_dic['fresh_ratio_schedule']
step = current['step']
num_steps = current['num_steps']
threshold = cache_dic['fresh_threshold']
weight = 0.9
if fresh_ratio_schedule == 'constant':
return fresh_ratio
elif fresh_ratio_schedule == 'linear':
return fresh_ratio * (1 + weight - 2 * weight * step / num_steps)
elif fresh_ratio_schedule == 'exp':
#return 0.5 * (0.052 ** (step/num_steps))
return fresh_ratio * (weight ** (step / num_steps))
elif fresh_ratio_schedule == 'linear-mode':
mode = (step % threshold)/threshold - 0.5
mode_weight = 0.1
return fresh_ratio * (1 + weight - 2 * weight * step / num_steps + mode_weight * mode)
elif fresh_ratio_schedule == 'layerwise':
return fresh_ratio * (1 + weight - 2 * weight * current['layer'] / 27)
elif fresh_ratio_schedule == 'linear-layerwise':
step_weight = -0.9 #0.9
step_factor = 1 - step_weight + 2 * step_weight * step / num_steps
#if current['layer'] == 2:
# return 1.0
#sigmoid
#sigmoid_weight = 0.13
#layer_factor = 2 * torch.sigmoid(torch.tensor([sigmoid_weight * (13.5 - current['layer'])]))
layer_weight = 0.6
layer_factor = 1 + layer_weight - 2 * layer_weight * current['layer'] / 27
module_weight = 1.0 #TokenCache N=8 2.5 N=6 2.5 #N=4 2.1
module_time_weight = 0.6
module_factor = (1 - (1-module_time_weight) * module_weight) if current['module']=='cross-attn' else (1 + module_time_weight * module_weight)
return fresh_ratio * layer_factor * step_factor * module_factor
elif fresh_ratio_schedule == 'ToCa':
step_weight = 0.0 #0.9
step_factor = 1 - step_weight + 2 * step_weight * step / num_steps
layer_weight = 0.5
layer_factor = 1 + layer_weight - 2 * layer_weight * current['layer'] / 27
#module_weight = 1.0
#module_time_weight = 0.6
# this means 60*x% cross-attn computation, and 160*x% mlp computation. This is designed for cross-attn has best temporal redundancy, and mlp has worse.
# so cross-attn compute less and mlp compute more.
#module_factor = (1 - (1-module_time_weight) * module_weight) if current['module']=='cross-attn' else (1 + module_time_weight * module_weight)
stream_weight = 0.6
stream_factor = (1 - stream_weight) if current['stream']=='double_stream' else (1 + stream_weight)
return fresh_ratio * layer_factor * step_factor * stream_factor #* module_factor
else:
raise ValueError("unrecognized fresh ratio schedule", fresh_ratio_schedule)
@@ -1,21 +0,0 @@
from .force_scheduler import force_scheduler
def global_force_fresh(cache_dic, current):
'''
Return whether to force fresh tokens globally.
'''
first_step = (current['step'] == 0)
second_step = (current['step'] == 1)
force_fresh = cache_dic['force_fresh']
if not first_step:
fresh_threshold = cache_dic['cal_threshold']
else:
fresh_threshold = cache_dic['fresh_threshold']
if force_fresh == 'global':
return (first_step or (current['step']% fresh_threshold == 0))
elif force_fresh == 'local':
return first_step
elif force_fresh == 'none':
return first_step
else:
raise ValueError("unrecognized force fresh strategy", force_fresh)
@@ -1,60 +0,0 @@
import torch
import torch.nn as nn
from .scores import attn_score, similarity_score, norm_score, k_norm_score, v_norm_score
def score_evaluate(cache_dic, tokens, current) -> torch.Tensor:
'''
Return the score tensor (B, N) for the given tokens.
'''
#if ((not current['is_force_fresh']) and (cache_dic['force_fresh'] == 'local')):
# # abandoned branch, if you want to explore the local force fresh strategy, this may help.
# force_fresh_mask = torch.as_tensor((cache_dic['cache_index'][-1][current['layer']][current['module']] >= 2 * cache_dic['fresh_threshold']), dtype = int) # 2 because the threshold is for step, not module
# force_len = force_fresh_mask.sum(dim=1)
# force_indices = force_fresh_mask.argsort(dim = -1, descending = True)[:, :force_len.min()]
# force_indices = force_indices[:, torch.randperm(force_indices.shape[1])]
# Just see more explanation in the version of DiT-ToCa if needed.
if cache_dic['cache_type'] == 'random':
score = torch.rand(tokens.shape[0], tokens.shape[1], device=tokens.device)
elif cache_dic['cache_type'] == 'straight':
score = torch.ones(tokens.shape[0], tokens.shape[1]).to(tokens.device)
elif cache_dic['cache_type'] == 'attention':
# cache_dic['attn_map'][step][layer] (B, N, N), the last dimention has get softmaxed
score = attn_score(cache_dic, current)
#score = score + 0.0 * torch.rand_like(score, device= score.device)
elif cache_dic['cache_type'] == 'similarity':
score = similarity_score(cache_dic, current, tokens)
elif cache_dic['cache_type'] == 'norm':
score = norm_score(cache_dic, current, tokens)
elif cache_dic['cache_type'] == 'k-norm':
score = k_norm_score(cache_dic, current)
elif cache_dic['cache_type'] == 'v-norm':
score = v_norm_score(cache_dic, current)
elif cache_dic['cache_type'] == 'compress':
score1 = torch.rand(int(tokens.shape[0]*0.5), tokens.shape[1])
score1 = torch.cat([score1, score1], dim=0).to(tokens.device)
score2 = cache_dic['attn_map'][-1][current['layer']].sum(dim=1)#.mean(dim=0) # (B, N)
# normalize
score2 = score2 / score2.max(dim=1, keepdim=True)[0]
score = 0.5 * score1 + 0.5 * score2
# abandoned the branch, if you want to explore the local force fresh strategy, this may help.
#if ((not current['is_force_fresh']) and (cache_dic['force_fresh'] == 'local')): # current['is_force_fresh'] is False, cause when it is True, no cut and fresh are needed
# #print(torch.ones_like(force_indices, dtype=float, device=force_indices.device).dtype)
# score.scatter_(dim=1, index=force_indices, src=torch.ones_like(force_indices, dtype=torch.float32,
# device=force_indices.device))
###if (True and (cache_dic['force_fresh'] == 'global')):
### soft_step_score = cache_dic['cache_index'][-1][current['layer']][current['module']].float() / (cache_dic['fresh_threshold'])
### #soft_layer_score = cache_dic['cache_index']['layer_index'][current['module']].float() / (27)
### score = score + cache_dic['soft_fresh_weight'] * soft_step_score #+ 0.1 *soft_layer_score
return score.to(tokens.device)
-77
View File
@@ -1,77 +0,0 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
def attn_score(cache_dic, current):
#self_attn_score = 1- cache_dic['attn_map'][-1][current['layer']].diagonal(dim1=1, dim2=2)
#self_attn_score = F.normalize(self_attn_score, dim=1, p=2)
#attention_score = F.normalize(cache_dic['attn_map'][-1][current['layer']].sum(dim=1), dim=1, p=2)
#cross_attn_map = F.threshold(cache_dic['cross_attn_map'][-1][current['layer']],threshold=0.0, value=0.0)
#cross_attention_score = F.normalize(cross_attn_map.sum(dim=-1), dim=-1, p=2)
# Note: It is important to give a same selection method for cfg and no cfg.
# Because the influence of **Cross-Attention** in text-contidional models makes cfg and no cfg a BIG difference.
# Same selection for cfg and no cfg
#cond_cmap, uncond_cmap = torch.split(cache_dic['attn_map'][-1][current['layer']], len(cache_dic['cross_attn_map'][-1][current['layer']]) // 2, dim=0)
#cond_weight = 0.5
#cmap = cond_weight * cond_cmap + (1 - cond_weight) * uncond_cmap
## Entropy score
#cross_attention_entropy = -torch.sum(cmap * torch.log(cmap + 1e-7), dim=-1)
#cross_attention_score = F.normalize(1 + cross_attention_entropy, dim=1, p=2) # Note here "1" does not influence the sorted sequence, but provie stability.
#score = cross_attention_score.repeat(2, 1)
if current['stream'] == 'double_stream':
score = F.normalize(cache_dic['attn_map'][-1][current['stream']][current['layer']][current['module']], dim=-1, p=2)
elif current['stream'] == 'single_stream':
score = F.normalize(cache_dic['attn_map'][-1][current['stream']][current['layer']]['total'], dim=-1, p=2)
# You can try conbining the self_attention_score (s1) and cross_attention_score (s2) as the final score, there exists a balance.
#cross_weight = 0.0
#score = (1-cross_weight) * attention_score + cross_weight * cross_attention_score
return score
def similarity_score(cache_dic, current, tokens):
cosine_sim = F.cosine_similarity(tokens, cache_dic['cache'][-1][current['layer']][current['module']], dim=-1)
return F.normalize(1- cosine_sim, dim=-1, p=2)
def norm_score(cache_dic, current, tokens):
norm = tokens.norm(dim=-1, p=2)
return F.normalize(norm, dim=-1, p=2)
def kv_norm_score(cache_dic, current):
# (B, N, num_heads)
#cond_k_norm, uncond_k_norm = torch.split(cache_dic['cache'][-1][current['layer']]['k_norm'], len(cache_dic['cache'][-1][current['layer']]['k_norm']) // 2, dim=0)
cond_v_norm, uncond_v_norm = torch.split(cache_dic['cache'][-1][current['layer']]['v_norm'], len(cache_dic['cache'][-1][current['layer']]['v_norm']) // 2, dim=0)
cond_weight = 0.5
#k_norm = cond_weight * cond_k_norm + (1 - cond_weight) * uncond_k_norm
v_norm = cond_weight * cond_v_norm + (1 - cond_weight) * uncond_v_norm
kv_norm = 1 -v_norm
## 计算 (B/2, N) 张量在 N 维度上的每个元素与均值的绝对值差
#kv_norm_mean = kv_norm.mean(dim=-2, keepdim=True)
#kv_norm_diff = torch.abs(kv_norm - kv_norm_mean)
return F.normalize(kv_norm.sum(dim=-1), p=2).repeat(2, 1)
def k_norm_score(cache_dic, current):
# (B, N)
if current['stream'] == 'double_stream':
score = F.normalize(cache_dic['k-norm'][-1][current['stream']][current['layer']][current['module']], dim=-1, p=2)
elif current['stream'] == 'single_stream':
score = F.normalize(cache_dic['k-norm'][-1][current['stream']][current['layer']]['total'], dim=-1, p=2)
return score
def v_norm_score(cache_dic, current):
# (B, N)
if current['stream'] == 'double_stream':
score = F.normalize(cache_dic['v-norm'][-1][current['stream']][current['layer']][current['module']], dim=-1, p=2)
elif current['stream'] == 'single_stream':
score = F.normalize(cache_dic['v-norm'][-1][current['stream']][current['layer']]['total'], dim=-1, p=2)
return score
@@ -1,52 +0,0 @@
import torch
from typing import Dict
def support_set_selection(x: torch.Tensor, fresh_ratio: float, base_ratio: float, current: Dict, cache_dic: Dict) -> torch.Tensor:
#selection_start = 0
#
#if current['stream'] == 'single_stream':
# # only select from the img tokens
# x = x[:, cache_dic['txt_shape'] :]
# selection_start = cache_dic['txt_shape']
B, N, H = x.shape
num_total = int(fresh_ratio * N) # 最终每个 batch 选取的 token 数
base_count = int(base_ratio * num_total) # 随机选取的 token 数
#base_count = 1
add_count = num_total - base_count # 需要从候选集中选取的 token 数
# 1. 随机选取 (B, base_count) 个 token
random_indices = torch.randperm(N, device=x.device)
base_indices = random_indices[:base_count]
other_indices = random_indices[base_count:]
base_tokens = x.gather(dim=1, index=base_indices.unsqueeze(-1).expand(B, -1, H))
#other_tokens = x.gather(dim=1, index=other_indices.unsqueeze(-1).expand(-1, -1, H))
# 2. 计算余下 token 与已选 token 的相似度
# normaize
base_tokens = base_tokens / base_tokens.norm(dim=-1, keepdim=True)
#other_tokens = other_tokens / other_tokens.norm(dim=-1, keepdim=True)
x_norm = x / x.norm(dim=-1, keepdim=True)
# 计算余下 token 与已选 token 的相似度
similarity = torch.einsum('bnd,bmd->bnm', base_tokens, x_norm)
# 计算每列最小值
min_similarity = similarity.min(dim=1).values
#min_similarity = similarity.max(dim=1).values
# 3. 选取相似度最小的 token
_, min_indices = min_similarity.topk(add_count, largest=False)
#_, min_indices = min_similarity.topk(add_count, largest=True)
# 4. 合并 base_indices 和 min_indices
#indices = torch.cat([base_indices, other_indices[min_indices]], dim=-1)
indices = torch.cat([base_indices.expand(B, -1), min_indices], dim=-1) #+ selection_start
return indices
@@ -1,28 +0,0 @@
import torch
def token_merge(cache_dic, tokens, current, fresh_indices, stale_indices):
'''
An abandoned branch in exploring if token merge helps. The answer is no, at least no for training-free strategy.
'''
if (current['layer'] % 1 == 0):
fresh_tokens = torch.gather(input = tokens, dim = 1, index = fresh_indices.unsqueeze(-1).expand(-1, -1, tokens.shape[-1]))
stale_tokens = torch.gather(input = tokens, dim = 1, index = stale_indices.unsqueeze(-1).expand(-1, -1, tokens.shape[-1]))
method = 'similarity'
if method == 'distance':
descending = False
distance = torch.cdist(stale_tokens, fresh_tokens, p=1)
stale_fresh_dist, stale_fresh_indices_allstale = torch.min(distance, dim=2)
elif method == 'similarity':
descending = True
fresh_tokens = torch.nn.functional.normalize(fresh_tokens, p=2, dim=-1)
stale_tokens = torch.nn.functional.normalize(stale_tokens, p=2, dim=-1)
similarity = stale_tokens @ fresh_tokens.transpose(1, 2)
stale_fresh_dist, stale_fresh_indices_allstale = torch.max(similarity, dim=2)
saved_topk_stale = int((stale_fresh_dist > 0.995).sum(dim=1).min())
merged_stale_sequence = torch.sort(stale_fresh_dist, dim=1, descending=descending)[1][:,:saved_topk_stale]
stale_fresh_indices = stale_fresh_indices_allstale.gather(1, merged_stale_sequence)
merged_stale_sequence = stale_indices.gather(1, merged_stale_sequence)
merged_stale_fresh_indices = fresh_indices.gather(1, stale_fresh_indices)
cache_dic['merged_stale_fresh_indices'] = merged_stale_fresh_indices
cache_dic['merged_stale_sequence'] = merged_stale_sequence
@@ -1,19 +0,0 @@
import torch
def update_cache(fresh_indices, fresh_tokens, cache_dic, current, fresh_attn_map=None):
'''
Update the cache with the fresh tokens.
'''
step = current['step']
layer = current['layer']
module = current['module']
# Update the cached tokens at the positions
indices = fresh_indices
cache_dic['cache'][-1][current['stream']][current['layer']][current['module']][0].scatter_(dim=1, index=indices.unsqueeze(-1).expand(-1, -1, fresh_tokens.shape[-1]), src=fresh_tokens)
+23 -20
View File
@@ -4,6 +4,7 @@ import torch.nn as nn
from torch.nn import functional as F from torch.nn import functional as F
from comfy.utils import load_torch_file from comfy.utils import load_torch_file
@torch.compiler.disable()
def get_fp_maxval(bits=8, mantissa_bit=3, sign_bits=1): def get_fp_maxval(bits=8, mantissa_bit=3, sign_bits=1):
_bits = torch.tensor(bits) _bits = torch.tensor(bits)
_mantissa_bit = torch.tensor(mantissa_bit) _mantissa_bit = torch.tensor(mantissa_bit)
@@ -17,6 +18,7 @@ def get_fp_maxval(bits=8, mantissa_bit=3, sign_bits=1):
maxval = mantissa * 2 ** (2**E - 1 - bias) maxval = mantissa * 2 ** (2**E - 1 - bias)
return maxval return maxval
@torch.compiler.disable()
def quantize_to_fp8(x, bits=8, mantissa_bit=3, sign_bits=1): def quantize_to_fp8(x, bits=8, mantissa_bit=3, sign_bits=1):
""" """
Default is E4M3. Default is E4M3.
@@ -40,6 +42,7 @@ def quantize_to_fp8(x, bits=8, mantissa_bit=3, sign_bits=1):
qdq_out = torch.round(input_clamp / log_scales) * log_scales qdq_out = torch.round(input_clamp / log_scales) * log_scales
return qdq_out, log_scales return qdq_out, log_scales
@torch.compiler.disable()
def fp8_tensor_quant(x, scale, bits=8, mantissa_bit=3, sign_bits=1): def fp8_tensor_quant(x, scale, bits=8, mantissa_bit=3, sign_bits=1):
for i in range(len(x.shape) - 1): for i in range(len(x.shape) - 1):
scale = scale.unsqueeze(-1) scale = scale.unsqueeze(-1)
@@ -47,10 +50,10 @@ def fp8_tensor_quant(x, scale, bits=8, mantissa_bit=3, sign_bits=1):
quant_dequant_x, log_scales = quantize_to_fp8(new_x, bits=bits, mantissa_bit=mantissa_bit, sign_bits=sign_bits) quant_dequant_x, log_scales = quantize_to_fp8(new_x, bits=bits, mantissa_bit=mantissa_bit, sign_bits=sign_bits)
return quant_dequant_x, scale, log_scales return quant_dequant_x, scale, log_scales
def fp8_activation_dequant(qdq_out, scale, dtype): @torch.compiler.disable()
def fp8_activation_dequant(qdq_out, dtype):
qdq_out = qdq_out.type(dtype) qdq_out = qdq_out.type(dtype)
quant_dequant_x = qdq_out * scale.to(dtype) return qdq_out
return quant_dequant_x
def fp8_linear_forward(cls, original_dtype, input): def fp8_linear_forward(cls, original_dtype, input):
weight_dtype = cls.weight.dtype weight_dtype = cls.weight.dtype
@@ -62,33 +65,33 @@ def fp8_linear_forward(cls, original_dtype, input):
linear_weight = linear_weight.to(torch.float8_e4m3fn) linear_weight = linear_weight.to(torch.float8_e4m3fn)
weight_dtype = linear_weight.dtype weight_dtype = linear_weight.dtype
else: else:
scale = cls.fp8_scale.to(cls.weight.device) scale = cls.fp8_scale#.to(cls.weight.device)
linear_weight = cls.weight linear_weight = cls.weight
##### #####
if weight_dtype == torch.float8_e4m3fn and cls.weight.sum() != 0: #if weight_dtype == torch.float8_e4m3fn and cls.weight.sum() != 0:
if True or len(input.shape) == 3: if weight_dtype == torch.float8_e4m3fn:
cls_dequant = fp8_activation_dequant(linear_weight, scale, original_dtype) qdq_out = fp8_activation_dequant(linear_weight, original_dtype)
if cls.bias != None: cls_dequant = qdq_out * scale
output = F.linear(input, cls_dequant, cls.bias) if cls.bias != None:
else: output = F.linear(input, cls_dequant, cls.bias)
output = F.linear(input, cls_dequant)
return output
else: else:
return cls.original_forward(input.to(original_dtype)) output = F.linear(input, cls_dequant)
return output
else: else:
return cls.original_forward(input) return cls.original_forward(input)
def convert_fp8_linear(module, original_dtype): def convert_fp8_linear(module, original_dtype, device, fp8_scale_map={}):
setattr(module, "fp8_matmul_enabled", True) setattr(module, "fp8_matmul_enabled", True)
script_directory = os.path.dirname(os.path.abspath(__file__)) script_directory = os.path.dirname(os.path.abspath(__file__))
# loading fp8 mapping file # loading fp8 mapping file
fp8_map_path = os.path.join(script_directory,"fp8_map.safetensors") if not fp8_scale_map:
if os.path.exists(fp8_map_path): fp8_map_path = os.path.join(script_directory,"fp8_map.safetensors")
fp8_map = load_torch_file(fp8_map_path, safe_load=True) if os.path.exists(fp8_map_path):
else: fp8_map = load_torch_file(fp8_map_path, safe_load=True)
raise ValueError(f"Invalid fp8_map path: {fp8_map_path}.") else:
raise ValueError(f"Invalid fp8_map path: {fp8_map_path}.")
#fp8_layers = [] #fp8_layers = []
for key, layer in module.named_modules(): for key, layer in module.named_modules():
@@ -96,6 +99,6 @@ def convert_fp8_linear(module, original_dtype):
#fp8_layers.append(key) #fp8_layers.append(key)
original_forward = layer.forward original_forward = layer.forward
#layer.weight = torch.nn.Parameter(layer.weight.to(torch.float8_e4m3fn)) #layer.weight = torch.nn.Parameter(layer.weight.to(torch.float8_e4m3fn))
setattr(layer, "fp8_scale", fp8_map[key].to(dtype=original_dtype)) setattr(layer, "fp8_scale", fp8_map[key].to(device=device, dtype=original_dtype))
setattr(layer, "original_forward", original_forward) setattr(layer, "original_forward", original_forward)
setattr(layer, "forward", lambda input, m=layer: fp8_linear_forward(m, original_dtype, input)) setattr(layer, "forward", lambda input, m=layer: fp8_linear_forward(m, original_dtype, input))
+240 -421
View File
@@ -21,9 +21,6 @@ from ...enhance_a_video.enhance import get_feta_scores
from ...enhance_a_video.globals import is_enhance_enabled_single, is_enhance_enabled_double, set_num_frames from ...enhance_a_video.globals import is_enhance_enabled_single, is_enhance_enabled_double, set_num_frames
from .norm_layers import RMSNorm from .norm_layers import RMSNorm
from .cache_functions import cal_type
from .taylor_utils import derivative_approximation, taylor_formula, taylor_cache_init
from contextlib import contextmanager from contextlib import contextmanager
@contextmanager @contextmanager
@@ -203,8 +200,6 @@ class MMDoubleStreamBlock(nn.Module):
token_replace_vec: torch.Tensor = None, token_replace_vec: torch.Tensor = None,
first_frame_token_num: int = None, first_frame_token_num: int = None,
condition_type: str = None, condition_type: str = None,
cache_dic: Optional[Dict] = None,
current: Optional[Dict] = None,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> Tuple[torch.Tensor, torch.Tensor]:
if condition_type == "token_replace": if condition_type == "token_replace":
img_mod1, token_replace_img_mod1 = self.img_mod(vec, condition_type=condition_type, \ img_mod1, token_replace_img_mod1 = self.img_mod(vec, condition_type=condition_type, \
@@ -241,245 +236,108 @@ class MMDoubleStreamBlock(nn.Module):
) = self.txt_mod(vec).chunk(6, dim=-1) ) = self.txt_mod(vec).chunk(6, dim=-1)
# Prepare image for attention. # Prepare image for attention.
if cache_dic is None: img_modulated = self.img_norm1(img)
img_modulated = self.img_norm1(img) if condition_type == "token_replace":
if condition_type == "token_replace": img_modulated = modulate(
img_modulated = modulate( img_modulated, shift=img_mod1_shift, scale=img_mod1_scale, condition_type=condition_type,
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale, condition_type=condition_type, tr_shift=tr_img_mod1_shift, tr_scale=tr_img_mod1_scale,
tr_shift=tr_img_mod1_shift, tr_scale=tr_img_mod1_scale, first_frame_token_num=first_frame_token_num
first_frame_token_num=first_frame_token_num
)
else:
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
)
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
) )
# Apply QK-Norm if needed else:
img_q = self.img_attn_q_norm(img_q).to(img_v) img_modulated = modulate(
img_k = self.img_attn_k_norm(img_k).to(img_v) img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
# Apply RoPE if needed.
if freqs_cis is not None:
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# Prepare txt for attention.
txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
) )
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply QK-Norm if needed. # Apply RoPE if needed.
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v) if freqs_cis is not None:
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v) img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
if is_enhance_enabled_double(): # Prepare txt for attention.
feta_scores = get_feta_scores(img_q, img_k) txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Run actual attention. # Apply QK-Norm if needed.
q = torch.cat((img_q, txt_q), dim=1) txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
k = torch.cat((img_k, txt_k), dim=1) txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
v = torch.cat((img_v, txt_v), dim=1)
attn = attention( if is_enhance_enabled_double():
q, feta_scores = get_feta_scores(img_q, img_k)
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=img_k.shape[0],
attn_mask=attn_mask
)
img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :] # Run actual attention.
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
v = torch.cat((img_v, txt_v), dim=1)
if is_enhance_enabled_double(): attn = attention(
img_attn *= feta_scores q,
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=img_k.shape[0],
attn_mask=attn_mask
)
# Calculate the img bloks. img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :]
if condition_type == "token_replace": if is_enhance_enabled_double():
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate, condition_type=condition_type, img_attn *= feta_scores
tr_gate=tr_img_mod1_gate, first_frame_token_num=first_frame_token_num)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale, condition_type=condition_type,
tr_shift=tr_img_mod2_shift, tr_scale=tr_img_mod2_scale, first_frame_token_num=first_frame_token_num
)
),
gate=img_mod2_gate, condition_type=condition_type,
tr_gate=tr_img_mod2_gate, first_frame_token_num=first_frame_token_num
)
else:
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
),
gate=img_mod2_gate,
)
# Calculate the txt bloks. # Calculate the img bloks.
txt = txt + apply_gate(self.txt_attn_proj(txt_attn), gate=txt_mod1_gate) if condition_type == "token_replace":
txt = txt + apply_gate( img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate, condition_type=condition_type,
self.txt_mlp( tr_gate=tr_img_mod1_gate, first_frame_token_num=first_frame_token_num)
img = img + apply_gate(
self.img_mlp(
modulate( modulate(
self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale, condition_type=condition_type,
tr_shift=tr_img_mod2_shift, tr_scale=tr_img_mod2_scale, first_frame_token_num=first_frame_token_num
) )
), ),
gate=txt_mod2_gate, gate=img_mod2_gate, condition_type=condition_type,
tr_gate=tr_img_mod2_gate, first_frame_token_num=first_frame_token_num
)
else:
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
),
gate=img_mod2_gate,
) )
return img, txt # Calculate the txt bloks.
else: txt = txt + apply_gate(self.txt_attn_proj(txt_attn), gate=txt_mod1_gate)
if current['type'] == 'full': txt = txt + apply_gate(
current['module'] = 'attn' self.txt_mlp(
modulate(
img_modulated = self.img_norm1(img) self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale
if condition_type == "token_replace":
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale, condition_type=condition_type,
tr_shift=tr_img_mod1_shift, tr_scale=tr_img_mod1_scale,
first_frame_token_num=first_frame_token_num
)
else:
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
)
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
) )
# Apply QK-Norm if needed ),
img_q = self.img_attn_q_norm(img_q).to(img_v) gate=txt_mod2_gate,
img_k = self.img_attn_k_norm(img_k).to(img_v) )
# Apply RoPE if needed. return img, txt
if freqs_cis is not None:
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# Prepare txt for attention.
txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed.
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
if is_enhance_enabled_double():
feta_scores = get_feta_scores(img_q, img_k)
# Run actual attention.
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
v = torch.cat((img_v, txt_v), dim=1)
attn = attention(
q,
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=img_k.shape[0],
attn_mask=attn_mask
)
img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :]
if is_enhance_enabled_double():
img_attn *= feta_scores
# Calculate the img blocks
current['module'] = 'img_attn'
taylor_cache_init(cache_dic, current)
if condition_type == "token_replace":
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate, condition_type=condition_type,
tr_gate=tr_img_mod1_gate, first_frame_token_num=first_frame_token_num)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale, condition_type=condition_type,
tr_shift=tr_img_mod2_shift, tr_scale=tr_img_mod2_scale, first_frame_token_num=first_frame_token_num
)
),
gate=img_mod2_gate, condition_type=condition_type,
tr_gate=tr_img_mod2_gate, first_frame_token_num=first_frame_token_num
)
else:
#img attn
img_attn_out = self.img_attn_proj(img_attn)
img = img + apply_gate(img_attn_out, gate=img_mod1_gate)
derivative_approximation(cache_dic, current, img_attn_out)
#img mlp
current['module'] = 'img_mlp'
taylor_cache_init(cache_dic, current)
img_mlp_out = self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
)
img = img + apply_gate(img_mlp_out, gate=img_mod2_gate)
derivative_approximation(cache_dic, current, img_mlp_out)
# Calculate the txt blocks
current['module'] = 'txt_attn'
taylor_cache_init(cache_dic, current)
txt_attn_out = self.txt_attn_proj(txt_attn)
txt = txt + apply_gate(txt_attn_out, gate=txt_mod1_gate)
derivative_approximation(cache_dic, current, txt_attn_out)
current['module'] = 'txt_mlp'
taylor_cache_init(cache_dic, current)
txt_mlp_out = self.txt_mlp(
modulate(
self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale
)
)
txt = txt + apply_gate(txt_mlp_out, gate=txt_mod2_gate)
derivative_approximation(cache_dic, current, txt_mlp_out)
elif current['type'] == 'taylor_cache':
current['module'] = 'img_attn'
img = img + apply_gate(taylor_formula(cache_dic, current), gate=img_mod1_gate)
current['module'] = 'img_mlp'
img = img + apply_gate(taylor_formula(cache_dic, current), gate=img_mod2_gate)
current['module'] = 'txt_attn'
txt = txt + apply_gate(taylor_formula(cache_dic, current), gate=txt_mod1_gate)
current['module'] = 'txt_mlp'
txt = txt + apply_gate(taylor_formula(cache_dic, current),gate=txt_mod2_gate)
return img, txt
#region single block
class MMSingleStreamBlock(nn.Module): class MMSingleStreamBlock(nn.Module):
""" """
A DiT block with parallel linear layers as described in A DiT block with parallel linear layers as described in
@@ -568,8 +426,6 @@ class MMSingleStreamBlock(nn.Module):
token_replace_vec: torch.Tensor = None, token_replace_vec: torch.Tensor = None,
first_frame_token_num: int = None, first_frame_token_num: int = None,
condition_type: str = None, condition_type: str = None,
cache_dic: Optional[Dict] = None,
current: Optional[Dict] = None,
stg_mode: Optional[str] = None, stg_mode: Optional[str] = None,
) -> torch.Tensor: ) -> torch.Tensor:
@@ -585,81 +441,59 @@ class MMSingleStreamBlock(nn.Module):
tr_mod_gate) = tr_mod.chunk(3, dim=-1) tr_mod_gate) = tr_mod.chunk(3, dim=-1)
else: else:
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1) mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
if cache_dic is None: if condition_type == "token_replace":
if condition_type == "token_replace": x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale, condition_type=condition_type,
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale, condition_type=condition_type, tr_shift=tr_mod_shift, tr_scale=tr_mod_scale, first_frame_token_num=first_frame_token_num)
tr_shift=tr_mod_shift, tr_scale=tr_mod_scale, first_frame_token_num=first_frame_token_num) else:
else: x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale) qkv, mlp = torch.split(
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
)
qkv, mlp = torch.split( q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num) # Apply QK-Norm if needed.
q = self.q_norm(q).to(v)
k = self.k_norm(k).to(v)
# Apply QK-Norm if needed. # Apply RoPE if needed.
q = self.q_norm(q).to(v) if freqs_cis is not None:
k = self.k_norm(k).to(v) img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# assert (
# img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
# ), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
# Apply RoPE if needed. if is_enhance_enabled_single():
if freqs_cis is not None: feta_scores = get_feta_scores(img_q, img_k)
img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# assert (
# img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
# ), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
if is_enhance_enabled_single(): # Compute attention.
feta_scores = get_feta_scores(img_q, img_k) #assert (
# cu_seqlens_q.shape[0] == 2 * x.shape[0] + 1
# Compute attention. #), f"cu_seqlens_q.shape:{cu_seqlens_q.shape}, x.shape[0]:{x.shape[0]}"
#assert ( if stg_mode is not None:
# cu_seqlens_q.shape[0] == 2 * x.shape[0] + 1 if stg_mode == "STG-A":
#), f"cu_seqlens_q.shape:{cu_seqlens_q.shape}, x.shape[0]:{x.shape[0]}" attn = attention(
if stg_mode is not None: q,
if stg_mode == "STG-A": k,
attn = attention( v,
q, heads = self.heads_num,
k, mode=self.attention_mode,
v, cu_seqlens_q=cu_seqlens_q,
heads = self.heads_num, cu_seqlens_kv=cu_seqlens_kv,
mode=self.attention_mode, max_seqlen_q=max_seqlen_q,
cu_seqlens_q=cu_seqlens_q, max_seqlen_kv=max_seqlen_kv,
cu_seqlens_kv=cu_seqlens_kv, batch_size=x.shape[0],
max_seqlen_q=max_seqlen_q, do_stg=True,
max_seqlen_kv=max_seqlen_kv, txt_len=txt_len,
batch_size=x.shape[0], attn_mask=attn_mask
do_stg=True, )
txt_len=txt_len, output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
attn_mask=attn_mask return x + apply_gate(output, gate=mod_gate)
) elif stg_mode == "STG-R":
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
return x + apply_gate(output, gate=mod_gate)
elif stg_mode == "STG-R":
attn = attention(
q,
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=x.shape[0],
attn_mask=attn_mask
)
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
output = apply_gate(output, gate=mod_gate)
batch_size = output.shape[0]
output[:batch_size-1, :, :] = 0
return x + output
else:
attn = attention( attn = attention(
q, q,
k, k,
@@ -673,88 +507,32 @@ class MMSingleStreamBlock(nn.Module):
batch_size=x.shape[0], batch_size=x.shape[0],
attn_mask=attn_mask attn_mask=attn_mask
) )
if is_enhance_enabled_single():
attn *= feta_scores
#attn[:, :-txt_len, :] *= feta_scores
# Compute activation in mlp stream, cat again and run second linear layer. # Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2)) output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
if condition_type == "token_replace": output = apply_gate(output, gate=mod_gate)
output = x + apply_gate(output, gate=mod_gate, condition_type=condition_type, batch_size = output.shape[0]
tr_gate=tr_mod_gate, first_frame_token_num=first_frame_token_num) output[:batch_size-1, :, :] = 0
return output return x + output
else:
return x + apply_gate(output, gate=mod_gate)
else: else:
if current['type'] == 'full': attn = attention(
q,
#current['module'] = 'mlp' k,
#taylor_cache_init(cache_dic, current) v,
heads = self.heads_num,
if condition_type == "token_replace": mode=self.attention_mode,
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale, condition_type=condition_type, cu_seqlens_q=cu_seqlens_q,
tr_shift=tr_mod_shift, tr_scale=tr_mod_scale, first_frame_token_num=first_frame_token_num) cu_seqlens_kv=cu_seqlens_kv,
else: max_seqlen_q=max_seqlen_q,
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale) max_seqlen_kv=max_seqlen_kv,
batch_size=x.shape[0],
qkv, mlp = torch.split( attn_mask=attn_mask
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1 )
) if is_enhance_enabled_single():
attn *= feta_scores
current['module'] = 'attn' #attn[:, :-txt_len, :] *= feta_scores
taylor_cache_init(cache_dic, current)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed.
q = self.q_norm(q).to(v)
k = self.k_norm(k).to(v)
# Apply RoPE if needed.
if freqs_cis is not None:
img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# assert (
# img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
# ), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
if is_enhance_enabled_single():
feta_scores = get_feta_scores(img_q, img_k)
# Compute attention.
attn = attention(
q,
k,
v,
heads=self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=x.shape[0],
attn_mask=attn_mask
)
if is_enhance_enabled_single():
attn *= feta_scores
#attn[:, :-txt_len, :] *= feta_scores
derivative_approximation(cache_dic, current, attn)
current['module'] = 'total'
taylor_cache_init(cache_dic, current)
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
derivative_approximation(cache_dic, current, output)
elif current['type'] == 'taylor_cache':
current['module'] = 'total'
output = taylor_formula(cache_dic, current)
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
if condition_type == "token_replace": if condition_type == "token_replace":
output = x + apply_gate(output, gate=mod_gate, condition_type=condition_type, output = x + apply_gate(output, gate=mod_gate, condition_type=condition_type,
tr_gate=tr_mod_gate, first_frame_token_num=first_frame_token_num) tr_gate=tr_mod_gate, first_frame_token_num=first_frame_token_num)
@@ -974,15 +752,24 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
self.enable_teacache = False self.enable_teacache = False
self.cnt = 0 self.cnt = 0
self.num_steps = 0 self.num_steps = 0
self.teacache_skipped_steps = 0 self.teacache_skipped_steps_cond = 0
self.teacache_skipped_steps_uncond = 0
self.teacache_start_step = 0
self.teacache_end_step = 100
self.rel_l1_thresh = 0.15 self.rel_l1_thresh = 0.15
self.accumulated_rel_l1_distance = 0 self.accumulated_rel_l1_distance = 0
self.previous_modulated_input = None self.previous_modulated_input_cond = None
self.previous_residual = None self.previous_modulated_input_uncond = None
self.last_dimensions = None self.previous_residual_cond = None
self.last_frame_count = None self.previous_residual_uncond = None
self.teacache_device = None self.teacache_device = None
#slg
self.slg_single_blocks = None
self.slg_double_blocks = None
self.slg_start_percent = 0.0
self.slg_end_percent = 1.0
# thanks @2kpr for the initial block swap code! # thanks @2kpr for the initial block swap code!
def block_swap(self, double_blocks_to_swap, single_blocks_to_swap, offload_txt_in=False, offload_img_in=False): def block_swap(self, double_blocks_to_swap, single_blocks_to_swap, offload_txt_in=False, offload_img_in=False):
print(f"Swapping {double_blocks_to_swap + 1} double blocks and {single_blocks_to_swap + 1} single blocks") print(f"Swapping {double_blocks_to_swap + 1} double blocks and {single_blocks_to_swap + 1} single blocks")
@@ -1172,32 +959,40 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
stg_mode: str = None, stg_mode: str = None,
stg_block_idx: int = -1, stg_block_idx: int = -1,
return_dict: bool = True, return_dict: bool = True,
tseercache_dict = None, ref_latents: torch.Tensor = None,
tseer_current = None, is_uncond = False,
current_step: int = 0,
current_step_percentage: float = 0,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]: ) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
def _process_double_blocks(img, txt, vec, block_args, tseercache_dict=None, tseer_current=None): def _process_double_blocks(img, txt, vec, block_args):
for b, block in enumerate(self.double_blocks): for b, block in enumerate(self.double_blocks):
if self.slg_double_blocks is not None:
if b in self.slg_double_blocks and is_uncond:
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
print(f"Skipping double block {b}")
continue
if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0: if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0:
block.to(self.main_device) block.to(self.main_device)
if tseer_current is not None: img, txt = block(img, txt, vec, *block_args)
tseer_current['layer'] = b
img, txt = block(img, txt, vec, *block_args, tseercache_dict, tseer_current)
if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0: if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0:
block.to(self.offload_device, non_blocking=True) block.to(self.offload_device, non_blocking=True)
return img, txt return img, txt
def _process_single_blocks(x, vec, txt_seq_len, block_args, tseercache_dict=None, tseer_current=None, stg_mode=None, stg_block_idx=None): def _process_single_blocks(x, vec, txt_seq_len, block_args, stg_mode=None, stg_block_idx=None):
for b, block in enumerate(self.single_blocks): for b, block in enumerate(self.single_blocks):
if self.slg_single_blocks is not None:
if b in self.slg_single_blocks and is_uncond:
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
print(f"Skipping single block {b}")
continue
if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0: if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0:
block.to(self.main_device) block.to(self.main_device)
curr_stg_mode = stg_mode if b == stg_block_idx else None curr_stg_mode = stg_mode if b == stg_block_idx else None
if tseer_current is not None: x = block(x, vec, txt_seq_len, *block_args, curr_stg_mode)
tseer_current['layer'] = b
x = block(x, vec, txt_seq_len, *block_args, tseercache_dict, tseer_current, curr_stg_mode)
if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0: if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0:
block.to(self.offload_device, non_blocking=True) block.to(self.offload_device, non_blocking=True)
@@ -1258,6 +1053,12 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
self.img_in.to(self.main_device) self.img_in.to(self.main_device)
img = self.img_in(img) img = self.img_in(img)
if ref_latents is not None:
ref_latents = self.img_in(ref_latents)
ref_length = ref_latents.shape[-2]
img = torch.cat([ref_latents, img], dim=-2) # t c
if self.text_projection == "linear": if self.text_projection == "linear":
txt = self.txt_in(txt) txt = self.txt_in(txt)
elif self.text_projection == "single_refiner": elif self.text_projection == "single_refiner":
@@ -1298,7 +1099,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
] ]
#tea_cache #tea_cache
if self.enable_teacache: if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step:
inp = img.clone() inp = img.clone()
vec_ = vec.clone() vec_ = vec.clone()
txt_ = txt.clone() txt_ = txt.clone()
@@ -1316,29 +1117,51 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
normed_inp, shift=img_mod1_shift, scale=img_mod1_scale normed_inp, shift=img_mod1_shift, scale=img_mod1_scale
) )
# Choose the appropriate cache based on whether this is a conditional or unconditional pass
previous_modulated_input = self.previous_modulated_input_uncond if is_uncond else self.previous_modulated_input_cond
previous_residual = self.previous_residual_uncond if is_uncond else self.previous_residual_cond
accumulated_rel_l1_distance = self.accumulated_rel_l1_distance_uncond if is_uncond else self.accumulated_rel_l1_distance_cond
if self.cnt == 0 or self.cnt == self.num_steps-1: if self.cnt == 0 or self.cnt == self.num_steps-1:
should_calc = True should_calc = True
self.accumulated_rel_l1_distance = 0 accumulated_rel_l1_distance = 0
self.previous_modulated_input = modulated_inp.clone() previous_modulated_input = modulated_inp.clone()
else: else:
coefficients = [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02] coefficients = [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02]
rescale_func = np.poly1d(coefficients) rescale_func = np.poly1d(coefficients)
self.accumulated_rel_l1_distance += rescale_func(((modulated_inp-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean()).cpu().item()) if previous_modulated_input is not None:
if self.accumulated_rel_l1_distance < self.rel_l1_thresh: accumulated_rel_l1_distance += rescale_func(((modulated_inp-previous_modulated_input).abs().mean() / previous_modulated_input.abs().mean()).cpu().item())
should_calc = False if accumulated_rel_l1_distance < self.rel_l1_thresh:
should_calc = False
else:
should_calc = True
accumulated_rel_l1_distance = 0
else: else:
should_calc = True should_calc = True
self.accumulated_rel_l1_distance = 0 accumulated_rel_l1_distance = 0
self.previous_modulated_input = modulated_inp.clone()
# Store back the appropriate values
if is_uncond:
self.previous_modulated_input_uncond = modulated_inp.clone()
self.accumulated_rel_l1_distance_uncond = accumulated_rel_l1_distance
else:
self.previous_modulated_input_cond = modulated_inp.clone()
self.accumulated_rel_l1_distance_cond = accumulated_rel_l1_distance
self.cnt += 1 self.cnt += 1
if self.cnt == self.num_steps: if self.cnt == self.num_steps:
self.cnt = 0 self.cnt = 0
if not should_calc and self.previous_residual is not None: if not should_calc and previous_residual is not None:
self.teacache_skipped_steps += 1 # Increment the appropriate skipped steps counter
if is_uncond:
self.teacache_skipped_steps_uncond += 1
else:
self.teacache_skipped_steps_cond += 1
# Verify tensor dimensions match before adding # Verify tensor dimensions match before adding
if img.shape == self.previous_residual.shape: if img.shape == previous_residual.shape:
img = img + self.previous_residual.to(img.device) img = img + previous_residual.to(img.device)
else: else:
should_calc = True # Force recalculation if dimensions don't match should_calc = True # Force recalculation if dimensions don't match
@@ -1351,28 +1174,24 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx) x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx)
img = x[:, :img_seq_len, ...] img = x[:, :img_seq_len, ...]
self.previous_residual = (img - ori_img).to(self.teacache_device) new_residual = (img - ori_img).to(self.teacache_device)
else:
# TaylorSeer
if tseercache_dict is not None:
cal_type(tseercache_dict, tseer_current)
tseer_current['compute'] = not (tseer_current['type'] == 'aggressive')
if tseer_current['compute']:
tseer_current['stream'] = 'double_stream'
img, txt = _process_double_blocks(img, txt, vec, block_args, tseercache_dict, tseer_current)
x = torch.cat((img, txt), 1)
tseer_current['stream'] = 'single_stream'
x = _process_single_blocks(x, vec, txt.shape[1], block_args, tseercache_dict=tseercache_dict, tseer_current=tseer_current,
stg_mode=stg_mode, stg_block_idx=stg_block_idx)
else:
x = tseercache_dict['aggressive_feature']
else:
img, txt = _process_double_blocks(img, txt, vec, block_args)
x = torch.cat((img, txt), 1)
x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode=stg_mode, stg_block_idx=stg_block_idx)
# Store the new residual in the appropriate cache
if is_uncond:
self.previous_residual_uncond = new_residual
else:
self.previous_residual_cond = new_residual
else:
# Pass through DiT blocks
img, txt = _process_double_blocks(img, txt, vec, block_args)
# Merge txt and img to pass through single stream blocks.
x = torch.cat((img, txt), 1)
x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx)
img = x[:, :img_seq_len, ...] img = x[:, :img_seq_len, ...]
if ref_latents is not None:
img = img[:, ref_length:]
# ---------------------------- Final layer ------------------------------ # ---------------------------- Final layer ------------------------------
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels) img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
+49
View File
@@ -302,3 +302,52 @@ def get_1d_rotary_pos_embed_riflex(
torch.ones_like(freqs), freqs torch.ones_like(freqs), freqs
) # complex64 # [S, D/2] ) # complex64 # [S, D/2]
return freqs_cis return freqs_cis
def get_nd_rotary_pos_embed_new(rope_dim_list, start, *args, theta=10000., use_real=False,
theta_rescale_factor: Union[float, List[float]]=1.0,
interpolation_factor: Union[float, List[float]]=1.0,
concat_dict = {'mode': 'timecat-w', 'bias': -1}, num_frames: int = 129, k: int = 0,
):
grid = get_meshgrid_nd(start, *args, dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
if len(concat_dict)<1:
pass
else:
if concat_dict['mode']=='timecat':
bias = grid[:,:1].clone()
bias[0] = concat_dict['bias']*torch.ones_like(bias[0])
grid = torch.cat([bias, grid], dim=1)
elif concat_dict['mode']=='timecat-w':
bias = grid[:,:1].clone()
bias[0] = concat_dict['bias']*torch.ones_like(bias[0])
bias[2] += start[-1] ## ref https://github.com/Yuanshi9815/OminiControl/blob/main/src/generate.py#L178
grid = torch.cat([bias, grid], dim=1)
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
assert len(theta_rescale_factor) == len(rope_dim_list), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
assert len(interpolation_factor) == len(rope_dim_list), "len(interpolation_factor) should equal to len(rope_dim_list)"
# use 1/ndim of dimensions to encode grid_axis
embs = []
for i in range(len(rope_dim_list)):
emb = get_1d_rotary_pos_embed(rope_dim_list[i], grid[i].reshape(-1), theta, use_real=use_real,
theta_rescale_factor=theta_rescale_factor[i],
interpolation_factor=interpolation_factor[i]) # 2 x [WHD, rope_dim_list[i]]
embs.append(emb)
if use_real:
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
return cos, sin
else:
emb = torch.cat(embs, dim=1) # (WHD, D/2)
return emb
-52
View File
@@ -1,52 +0,0 @@
from typing import Dict
import torch
import math
@torch.compiler.disable()
def derivative_approximation(cache_dic: Dict, current: Dict, feature: torch.Tensor):
"""
Compute derivative approximation
:param cache_dic: Cache dictionary
:param current: Information of the current step
"""
difference_distance = current['activated_steps'][-1] - current['activated_steps'][-2]
updated_taylor_factors = {}
updated_taylor_factors[0] = feature.to(cache_dic['cache_device'])
for i in range(cache_dic['max_order']):
if (cache_dic['cache'][-1][current['stream']][current['layer']][current['module']].get(i, None) is not None) and (current['step'] > cache_dic['first_enhance'] - 2):
updated_factor = updated_taylor_factors[i].to(cache_dic['compute_device'])
cached = cache_dic['cache'][-1][current['stream']][current['layer']][current['module']][i].to(cache_dic['compute_device'], non_blocking=True)
updated_taylor_factors[i + 1] = ((updated_factor - cached) / difference_distance).to(cache_dic['cache_device'], non_blocking=True)
else:
break
cache_dic['cache'][-1][current['stream']][current['layer']][current['module']] = updated_taylor_factors
@torch.compiler.disable()
def taylor_formula(cache_dic: Dict, current: Dict) -> torch.Tensor:
"""
Compute Taylor expansion error
:param cache_dic: Cache dictionary
:param current: Information of the current step
"""
x = current['step'] - current['activated_steps'][-1]
#x = current['t'] - current['activated_times'][-1]
output = 0
for i in range(len(cache_dic['cache'][-1][current['stream']][current['layer']][current['module']])):
cached = cache_dic['cache'][-1][current['stream']][current['layer']][current['module']][i]
cached = cached.to(cache_dic['compute_device'], non_blocking=True)
output = output + (1 / math.factorial(i)) * cached * (x ** i)
return output
@torch.compiler.disable()
def taylor_cache_init(cache_dic: Dict, current: Dict):
"""
Initialize Taylor cache, expanding storage areas for Taylor series derivatives
:param cache_dic: Cache dictionary
:param current: Information of the current step
"""
if current['step'] == 0:
cache_dic['cache'][-1][current['stream']][current['layer']][current['module']] = {}
+222 -208
View File
@@ -2,7 +2,8 @@ import os
import torch import torch
import json import json
import gc import gc
from .utils import log, print_memory from tqdm import tqdm
from .utils import log, print_memory, optimized_scale
from diffusers.video_processor import VideoProcessor from diffusers.video_processor import VideoProcessor
from typing import List, Dict, Any, Tuple from typing import List, Dict, Any, Tuple
import numpy as np import numpy as np
@@ -220,6 +221,8 @@ class HyVideoTeaCache:
"rel_l1_thresh": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.01, "rel_l1_thresh": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.01,
"tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts"}), "tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts"}),
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}), "cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
"start_step": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1, "tooltip": "Start step to apply TeaCache"}),
"end_step": ("INT", {"default": -1, "min": -1, "max": 100, "step": 1, "tooltip": "End step to apply TeaCache"}),
}, },
} }
@@ -229,14 +232,16 @@ class HyVideoTeaCache:
CATEGORY = "HunyuanVideoWrapper" CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "TeaCache settings for HunyuanVideo to speed up inference" DESCRIPTION = "TeaCache settings for HunyuanVideo to speed up inference"
def process(self, rel_l1_thresh, cache_device): def process(self, rel_l1_thresh, cache_device, start_step, end_step):
if cache_device == "main_device": if cache_device == "main_device":
teacache_device = mm.get_torch_device() teacache_device = mm.get_torch_device()
else: else:
teacache_device = mm.unet_offload_device() teacache_device = mm.unet_offload_device()
teacache_args = { teacache_args = {
"rel_l1_thresh": rel_l1_thresh, "rel_l1_thresh": rel_l1_thresh,
"cache_device": teacache_device "cache_device": teacache_device,
"start_step": start_step,
"end_step": end_step
} }
return (teacache_args,) return (teacache_args,)
@@ -276,7 +281,7 @@ class HyVideoModelLoader:
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}), "model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
"base_precision": (["fp32", "bf16"], {"default": "bf16"}), "base_precision": (["fp32", "bf16"], {"default": "bf16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_scaled', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6", "torchao_int4", "torchao_int8"], {"default": 'disabled', "tooltip": "optional quantization method"}), "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_scaled'], {"default": 'disabled', "tooltip": "optional quantization method"}),
"load_device": (["main_device", "offload_device"], {"default": "main_device"}), "load_device": (["main_device", "offload_device"], {"default": "main_device"}),
}, },
"optional": { "optional": {
@@ -325,6 +330,8 @@ class HyVideoModelLoader:
in_channels = sd["img_in.proj.weight"].shape[1] in_channels = sd["img_in.proj.weight"].shape[1]
if in_channels == 16 and "i2v" in model.lower(): if in_channels == 16 and "i2v" in model.lower():
i2v_condition_type = "token_replace" i2v_condition_type = "token_replace"
elif in_channels == 16 and "custom" in model.lower():
i2v_condition_type = "reference"
else: else:
i2v_condition_type = "latent_concat" i2v_condition_type = "latent_concat"
log.info(f"Condition type: {i2v_condition_type}") log.info(f"Condition type: {i2v_condition_type}")
@@ -380,167 +387,98 @@ class HyVideoModelLoader:
comfy_model=comfy_model, comfy_model=comfy_model,
) )
if not "torchao" in quantization: log.info("Using accelerate to load and assign model weights to device...")
log.info("Using accelerate to load and assign model weights to device...") if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast" or quantization == "fp8_scaled":
if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast" or quantization == "fp8_scaled": fp8_scale_map = {}
dtype = torch.float8_e4m3fn if "fp8_scale" in sd:
elif quantization == "fp8_e5m2": for k, v in sd.items():
dtype = torch.float8_e5m2 if k.endswith(".fp8_scale"):
else: fp8_scale_map[k] = v
dtype = base_dtype dtype = torch.float8_e4m3fn
params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"} elif quantization == "fp8_e5m2":
for name, param in transformer.named_parameters(): dtype = torch.float8_e5m2
#print("Assigning Parameter name: ", name) else:
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype dtype = base_dtype
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"}
param_count = sum(1 for _ in transformer.named_parameters())
for name, param in tqdm(transformer.named_parameters(),
desc=f"Loading transformer parameters to {transformer_load_device}",
total=param_count,
leave=True):
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
comfy_model.diffusion_model = transformer comfy_model.diffusion_model = transformer
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
pipe.comfy_model = patcher pipe.comfy_model = patcher
del sd del sd
gc.collect() gc.collect()
mm.soft_empty_cache() mm.soft_empty_cache()
if lora is not None: if lora is not None:
from comfy.sd import load_lora_for_models from comfy.sd import load_lora_for_models
for l in lora: for l in lora:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}") log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
lora_path = l["path"] lora_path = l["path"]
lora_strength = l["strength"] lora_strength = l["strength"]
lora_sd = load_torch_file(lora_path, safe_load=True) lora_sd = load_torch_file(lora_path, safe_load=True)
lora_sd = standardize_lora_key_format(lora_sd) lora_sd = standardize_lora_key_format(lora_sd)
if l["blocks"]: if l["blocks"]:
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"]) lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"])
# patch in channels for keyframe LoRA # patch in channels for keyframe LoRA
if "diffusion_model.img_in.proj.lora_A.weight" in lora_sd: if "diffusion_model.img_in.proj.lora_A.weight" in lora_sd:
from .hyvideo.modules.embed_layers import PatchEmbed from .hyvideo.modules.embed_layers import PatchEmbed
if lora_sd["diffusion_model.img_in.proj.lora_A.weight"].shape[1] != in_channels: if lora_sd["diffusion_model.img_in.proj.lora_A.weight"].shape[1] != in_channels:
log.info(f"Different in_channels {lora_sd['diffusion_model.img_in.proj.lora_A.weight'].shape[1]} vs {in_channels}, patching...") log.info(f"Different in_channels {lora_sd['diffusion_model.img_in.proj.lora_A.weight'].shape[1]} vs {in_channels}, patching...")
new_img_in = PatchEmbed( new_img_in = PatchEmbed(
patch_size=patcher.model.diffusion_model.patch_size, patch_size=patcher.model.diffusion_model.patch_size,
in_chans=32, in_chans=32,
embed_dim=patcher.model.diffusion_model.hidden_size, embed_dim=patcher.model.diffusion_model.hidden_size,
).to(patcher.model.diffusion_model.device, dtype=patcher.model.diffusion_model.dtype) ).to(patcher.model.diffusion_model.device, dtype=patcher.model.diffusion_model.dtype)
new_img_in.proj.weight.zero_() new_img_in.proj.weight.zero_()
new_img_in.proj.weight[:, :in_channels].copy_(patcher.model.diffusion_model.img_in.proj.weight) new_img_in.proj.weight[:, :in_channels].copy_(patcher.model.diffusion_model.img_in.proj.weight)
if patcher.model.diffusion_model.img_in.proj.bias is not None: if patcher.model.diffusion_model.img_in.proj.bias is not None:
new_img_in.proj.bias.copy_(patcher.model.diffusion_model.img_in.proj.bias) new_img_in.proj.bias.copy_(patcher.model.diffusion_model.img_in.proj.bias)
patcher.model.diffusion_model.img_in = new_img_in patcher.model.diffusion_model.img_in = new_img_in
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0) patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
comfy.model_management.load_models_gpu([patcher]) comfy.model_management.load_models_gpu([patcher])
if load_device == "offload_device": if load_device == "offload_device":
patcher.model.diffusion_model.to(offload_device) patcher.model.diffusion_model.to(offload_device)
if quantization == "fp8_e4m3fn_fast": if quantization == "fp8_e4m3fn_fast":
from .fp8_optimization import convert_fp8_linear from .fp8_optimization import convert_fp8_linear
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep) params_to_keep.update({"mlp", "modulation", "mod"})
elif quantization == "fp8_scaled": convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep)
from .hyvideo.modules.fp8_optimization import convert_fp8_linear elif quantization == "fp8_scaled":
convert_fp8_linear(patcher.model.diffusion_model, base_dtype) from .hyvideo.modules.fp8_optimization import convert_fp8_linear
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, device, fp8_scale_map=fp8_scale_map)
if auto_cpu_offload: if auto_cpu_offload:
transformer.enable_auto_offload(dtype=dtype, device=device) if quantization == "fp8_scaled":
raise ValueError("Auto CPU offload and fp8 scaled quantization are not compatible.")
transformer.enable_auto_offload(dtype=dtype, device=device)
#compile #compile
if compile_args is not None: if compile_args is not None:
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"] if compile_args["compile_single_blocks"]:
if compile_args["compile_single_blocks"]: for i, block in enumerate(patcher.model.diffusion_model.single_blocks):
for i, block in enumerate(patcher.model.diffusion_model.single_blocks):
patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if compile_args["compile_double_blocks"]:
for i, block in enumerate(patcher.model.diffusion_model.double_blocks):
patcher.model.diffusion_model.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if compile_args["compile_txt_in"]:
patcher.model.diffusion_model.txt_in = torch.compile(patcher.model.diffusion_model.txt_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if compile_args["compile_vector_in"]:
patcher.model.diffusion_model.vector_in = torch.compile(patcher.model.diffusion_model.vector_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if compile_args["compile_final_layer"]:
patcher.model.diffusion_model.final_layer = torch.compile(patcher.model.diffusion_model.final_layer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
elif "torchao" in quantization:
try:
from torchao.quantization import (
quantize_,
fpx_weight_only,
float8_dynamic_activation_float8_weight,
int8_dynamic_activation_int8_weight,
int8_weight_only,
int4_weight_only
)
except:
raise ImportError("torchao is not installed")
# def filter_fn(module: nn.Module, fqn: str) -> bool:
# target_submodules = {'attn1', 'ff'} # avoid norm layers, 1.5 at least won't work with quantized norm1 #todo: test other models
# if any(sub in fqn for sub in target_submodules):
# return isinstance(module, nn.Linear)
# return False
if "fp6" in quantization:
quant_func = fpx_weight_only(3, 2)
elif "int4" in quantization:
quant_func = int4_weight_only()
elif "int8" in quantization:
quant_func = int8_weight_only()
elif "fp8dq" in quantization:
quant_func = float8_dynamic_activation_float8_weight()
elif 'fp8dqrow' in quantization:
from torchao.quantization.quant_api import PerRow
quant_func = float8_dynamic_activation_float8_weight(granularity=PerRow())
elif 'int8dq' in quantization:
quant_func = int8_dynamic_activation_int8_weight()
log.info(f"Quantizing model with {quant_func}")
comfy_model.diffusion_model = transformer
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
if lora is not None:
from comfy.sd import load_lora_for_models
for l in lora:
lora_path = l["path"]
lora_strength = l["strength"]
lora_sd = load_torch_file(lora_path, safe_load=True)
lora_sd = standardize_lora_key_format(lora_sd)
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
comfy.model_management.load_models_gpu([patcher])
for i, block in enumerate(patcher.model.diffusion_model.single_blocks):
log.info(f"Quantizing single_block {i}")
for name, _ in block.named_parameters(prefix=f"single_blocks.{i}"):
#print(f"Parameter name: {name}")
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
if compile_args is not None:
patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
quantize_(block, quant_func) if compile_args["compile_double_blocks"]:
print(block) for i, block in enumerate(patcher.model.diffusion_model.double_blocks):
block.to(offload_device)
for i, block in enumerate(patcher.model.diffusion_model.double_blocks):
log.info(f"Quantizing double_block {i}")
for name, _ in block.named_parameters(prefix=f"double_blocks.{i}"):
#print(f"Parameter name: {name}")
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
if compile_args is not None:
patcher.model.diffusion_model.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) patcher.model.diffusion_model.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
quantize_(block, quant_func) if compile_args["compile_txt_in"]:
for name, param in patcher.model.diffusion_model.named_parameters(): patcher.model.diffusion_model.txt_in = torch.compile(patcher.model.diffusion_model.txt_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if "single_blocks" not in name and "double_blocks" not in name: if compile_args["compile_vector_in"]:
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name]) patcher.model.diffusion_model.vector_in = torch.compile(patcher.model.diffusion_model.vector_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if compile_args["compile_final_layer"]:
manual_offloading = False # to disable manual .to(device) calls patcher.model.diffusion_model.final_layer = torch.compile(patcher.model.diffusion_model.final_layer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
log.info(f"Quantized transformer blocks to {quantization}")
for name, param in patcher.model.diffusion_model.named_parameters():
print(name, param.dtype)
#param.data = param.data.to(self.vae_dtype).to(device)
del sd
mm.soft_empty_cache()
patcher.model["pipe"] = pipe patcher.model["pipe"] = pipe
patcher.model["dtype"] = base_dtype patcher.model["dtype"] = base_dtype
@@ -629,10 +567,8 @@ class HyVideoTorchCompileSettings:
"compile_txt_in": ("BOOLEAN", {"default": False, "tooltip": "Compile txt_in layers"}), "compile_txt_in": ("BOOLEAN", {"default": False, "tooltip": "Compile txt_in layers"}),
"compile_vector_in": ("BOOLEAN", {"default": False, "tooltip": "Compile vector_in layers"}), "compile_vector_in": ("BOOLEAN", {"default": False, "tooltip": "Compile vector_in layers"}),
"compile_final_layer": ("BOOLEAN", {"default": False, "tooltip": "Compile final layer"}), "compile_final_layer": ("BOOLEAN", {"default": False, "tooltip": "Compile final layer"}),
}, },
"optional": {
"dynamo_recompile_limit": ("INT", {"default": 64, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.recompile_limit"}),
}
} }
RETURN_TYPES = ("COMPILEARGS",) RETURN_TYPES = ("COMPILEARGS",)
RETURN_NAMES = ("torch_compile_args",) RETURN_NAMES = ("torch_compile_args",)
@@ -640,7 +576,7 @@ class HyVideoTorchCompileSettings:
CATEGORY = "HunyuanVideoWrapper" CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch 2.5.0 is recommended" DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch 2.5.0 is recommended"
def loadmodel(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_single_blocks, compile_double_blocks, compile_txt_in, compile_vector_in, compile_final_layer, dynamo_recompile_limit=64): def loadmodel(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_single_blocks, compile_double_blocks, compile_txt_in, compile_vector_in, compile_final_layer):
compile_args = { compile_args = {
"backend": backend, "backend": backend,
@@ -648,7 +584,6 @@ class HyVideoTorchCompileSettings:
"mode": mode, "mode": mode,
"dynamic": dynamic, "dynamic": dynamic,
"dynamo_cache_size_limit": dynamo_cache_size_limit, "dynamo_cache_size_limit": dynamo_cache_size_limit,
"dynamo_recompile_limit": dynamo_recompile_limit,
"compile_single_blocks": compile_single_blocks, "compile_single_blocks": compile_single_blocks,
"compile_double_blocks": compile_double_blocks, "compile_double_blocks": compile_double_blocks,
"compile_txt_in": compile_txt_in, "compile_txt_in": compile_txt_in,
@@ -665,10 +600,14 @@ class HyVideoTextEmbedBridge:
def INPUT_TYPES(s): def INPUT_TYPES(s):
return {"required": { return {"required": {
"positive": ("CONDITIONING", ), "positive": ("CONDITIONING", ),
"cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "guidance scale"} ),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ),
"batched_cfg": ("BOOLEAN", {"default": False, "tooltip": "Calculate cond and uncond as a batch, increases memory usage but can be faster"}),
"use_cfg_zero_star": ("BOOLEAN", {"default": True, "tooltip": "Use CFG zero star"}),
}, },
"optional": { "optional": {
"negative": ("CONDITIONING", ), "negative": ("CONDITIONING", ),
"hyvid_cfg": ("HYVID_CFG", {"tooltip": "The prompt from the cfg node is not used, only the settings"}),
} }
} }
RETURN_TYPES = ("HYVIDEMBEDS",) RETURN_TYPES = ("HYVIDEMBEDS",)
@@ -677,7 +616,7 @@ class HyVideoTextEmbedBridge:
CATEGORY = "HunyuanVideoWrapper" CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "Acts as a bridge between the native ComfyUI conditioning and the HunyuanVideoWrapper embeds" DESCRIPTION = "Acts as a bridge between the native ComfyUI conditioning and the HunyuanVideoWrapper embeds"
def convert(self, positive, negative=None, hyvid_cfg=None): def convert(self, positive, cfg, start_percent, end_percent, batched_cfg, use_cfg_zero_star, negative=None):
positive_cond = positive[0][0] positive_cond = positive[0][0]
positive_pooled = positive[0][1]["pooled_output"] positive_pooled = positive[0][1]["pooled_output"]
positive_attention_mask = torch.ones(positive_cond.shape[1], dtype=torch.bool, device=positive_cond.device).unsqueeze(0) positive_attention_mask = torch.ones(positive_cond.shape[1], dtype=torch.bool, device=positive_cond.device).unsqueeze(0)
@@ -693,10 +632,11 @@ class HyVideoTextEmbedBridge:
"negative_attention_mask": negative_attention_mask, "negative_attention_mask": negative_attention_mask,
"prompt_embeds_2": positive_pooled, "prompt_embeds_2": positive_pooled,
"negative_prompt_embeds_2": negative_pooled, "negative_prompt_embeds_2": negative_pooled,
"cfg": torch.tensor(hyvid_cfg["cfg"]) if hyvid_cfg is not None else None, "cfg": torch.tensor(cfg),
"start_percent": torch.tensor(hyvid_cfg["start_percent"]) if hyvid_cfg is not None else None, "start_percent": torch.tensor(start_percent),
"end_percent": torch.tensor(hyvid_cfg["end_percent"]) if hyvid_cfg is not None else None, "end_percent": torch.tensor(end_percent),
"batched_cfg": torch.tensor(hyvid_cfg["batched_cfg"]) if hyvid_cfg is not None else None, "batched_cfg": torch.tensor(batched_cfg),
"use_cfg_zero_star": torch.tensor(use_cfg_zero_star),
} }
return (prompt_embeds_dict,) return (prompt_embeds_dict,)
@@ -1143,7 +1083,8 @@ class HyVideoCFG:
"cfg": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "guidance scale"} ), "cfg": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "guidance scale"} ),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ), "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ), "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ),
"batched_cfg": ("BOOLEAN", {"default": True, "tooltip": "Calculate cond and uncond as a batch, increases memory usage but can be faster"}), "batched_cfg": ("BOOLEAN", {"default": False, "tooltip": "Calculate cond and uncond as a batch, increases memory usage but can be faster"}),
"use_cfg_zero_star": ("BOOLEAN", {"default": False, "tooltip": "Use CFG zero star"}),
}, },
} }
@@ -1153,13 +1094,14 @@ class HyVideoCFG:
CATEGORY = "HunyuanVideoWrapper" CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "To use CFG with HunyuanVideo" DESCRIPTION = "To use CFG with HunyuanVideo"
def process(self, negative_prompt, cfg, start_percent, end_percent, batched_cfg): def process(self, negative_prompt, cfg, start_percent, end_percent, batched_cfg, use_cfg_zero_star):
cfg_dict = { cfg_dict = {
"negative_prompt": negative_prompt, "negative_prompt": negative_prompt,
"cfg": cfg, "cfg": cfg,
"start_percent": start_percent, "start_percent": start_percent,
"end_percent": end_percent, "end_percent": end_percent,
"batched_cfg": batched_cfg "batched_cfg": batched_cfg,
"use_cfg_zero_start": use_cfg_zero_star,
} }
return (cfg_dict,) return (cfg_dict,)
@@ -1238,6 +1180,7 @@ class HyVideoTextEmbedsLoad:
"start_percent": loaded_tensors.get("start_percent", None), "start_percent": loaded_tensors.get("start_percent", None),
"end_percent": loaded_tensors.get("end_percent", None), "end_percent": loaded_tensors.get("end_percent", None),
"batched_cfg": loaded_tensors.get("batched_cfg", None), "batched_cfg": loaded_tensors.get("batched_cfg", None),
"use_cfg_zero_star": loaded_tensors.get("use_cfg_zero_star", None),
} }
return (prompt_embeds_dict,) return (prompt_embeds_dict,)
@@ -1271,34 +1214,73 @@ class HyVideoContextOptions:
return (context_options,) return (context_options,)
class HyVideoTaylorSeerOptions: class HyVideoLoopArgs:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return {"required": { return {"required": {
"cache_device": (["main_device", "offload_device"], {"default": "offload_device"}), "shift_skip": ("INT", {"default": 6, "min": 0, "tooltip": "Skip step of latent shift"}),
"max_order": ("INT", {"default": 1, "min": 0, "max": 10, "step": 1, "tooltip": "Maximum order of the Taylor series expansion"}), "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the looping effect"}),
"fresh_threshold": ("INT", {"default": 5, "min": 0, "max": 100, "step": 1, "tooltip": "A higher fresh_threshold results in faster inference but may reduce generation quality."}), "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the looping effect"}),
} },
} }
RETURN_TYPES = ("TAYLORSEERARGS", ) RETURN_TYPES = ("LOOPARGS", )
RETURN_NAMES = ("taylorseer_args",) RETURN_NAMES = ("loop_args",)
FUNCTION = "passargs" FUNCTION = "process"
CATEGORY = "HunyuanVideoWrapper" CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "https://github.com/Shenyi-Z/TaylorSeer" DESCRIPTION = "Looping through latent shift as shown in https://github.com/YisuiTT/Mobius/"
def passargs(self, cache_device, max_order, fresh_threshold): def process(self, **kwargs):
if cache_device == "main_device": return (kwargs,)
cache_device = mm.get_torch_device()
else: class HunyuanVideoFresca:
cache_device = mm.unet_offload_device() @classmethod
args = { def INPUT_TYPES(s):
"cache_device": cache_device, return {"required": {
"compute_device": mm.get_torch_device(), "fresca_scale_low": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
"max_order": max_order, "fresca_scale_high": ("FLOAT", {"default": 1.25, "min": 0.0, "max": 10.0, "step": 0.01}),
"fresh_threshold": fresh_threshold, "fresca_freq_cutoff": ("INT", {"default": 20, "min": 0, "max": 10000, "step": 1}),
},
} }
return (args,)
RETURN_TYPES = ("FRESCA_ARGS", )
RETURN_NAMES = ("fresca_args",)
FUNCTION = "process"
CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "https://github.com/WikiChao/FreSca"
def process(self, **kwargs):
return (kwargs,)
class HunyuanVideoSLG:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"double_blocks": ("STRING", {"default": "", "tooltip": "Blocks to skip uncond on, separated by comma, index starts from 0"}),
"single_blocks": ("STRING", {"default": "20", "tooltip": "Blocks to skip uncond on, separated by comma, index starts from 0"}),
"start_percent": ("FLOAT", {"default": 0.4, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of SLG signal"}),
"end_percent": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of SLG signal"}),
},
}
RETURN_TYPES = ("SLGARGS", )
RETURN_NAMES = ("slg_args",)
FUNCTION = "process"
CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "Skips uncond on the selected blocks"
def process(self, double_blocks, single_blocks, start_percent, end_percent):
slg_double_block_list = [int(x.strip()) for x in double_blocks.split(",")] if double_blocks else None
slg_single_block_list = [int(x.strip()) for x in single_blocks.split(",")] if single_blocks else None
slg_args = {
"double_blocks": slg_double_block_list,
"single_blocks": slg_single_block_list,
"start_percent": start_percent,
"end_percent": end_percent,
}
return (slg_args,)
#region Sampler #region Sampler
class HyVideoSampler: class HyVideoSampler:
@@ -1321,6 +1303,7 @@ class HyVideoSampler:
"optional": { "optional": {
"samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ), "samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ),
"image_cond_latents": ("LATENT", {"tooltip": "init Latents to use for image2video process"} ), "image_cond_latents": ("LATENT", {"tooltip": "init Latents to use for image2video process"} ),
#"neg_image_cond_latents": ("LATENT", {"tooltip": "init Latents to use for image2video process"} ),
"denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), "denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"stg_args": ("STGARGS", ), "stg_args": ("STGARGS", ),
"context_options": ("HYVIDCONTEXT", ), "context_options": ("HYVIDCONTEXT", ),
@@ -1332,7 +1315,10 @@ 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"}),
"taylorseer_args": ("TAYLORSEERARGS", ), "loop_args": ("LOOPARGS", ),
"fresca_args": ("FRESCA_ARGS", ),
"slg_args": ("SLGARGS", ),
"mask": ("MASK", ),
} }
} }
@@ -1343,7 +1329,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, riflex_freq_index=0, i2v_mode="stability", taylorseer_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, fresca_args=None, slg_args=None, mask=None):
model = model.model model = model.model
device = mm.get_torch_device() device = mm.get_torch_device()
@@ -1366,11 +1352,13 @@ class HyVideoSampler:
cfg_start_percent = float(hyvid_embeds.get("start_percent", 0.0)) cfg_start_percent = float(hyvid_embeds.get("start_percent", 0.0))
cfg_end_percent = float(hyvid_embeds.get("end_percent", 1.0)) cfg_end_percent = float(hyvid_embeds.get("end_percent", 1.0))
batched_cfg = hyvid_embeds.get("batched_cfg", True) batched_cfg = hyvid_embeds.get("batched_cfg", True)
use_cfg_zero_star = hyvid_embeds.get("use_cfg_zero_star", True)
else: else:
cfg = 1.0 cfg = 1.0
cfg_start_percent = 0.0 cfg_start_percent = 0.0
cfg_end_percent = 1.0 cfg_end_percent = 1.0
batched_cfg = False batched_cfg = False
use_cfg_zero_star = False
if embedded_guidance_scale == 0.0: if embedded_guidance_scale == 0.0:
embedded_guidance_scale = None embedded_guidance_scale = None
@@ -1431,24 +1419,21 @@ class HyVideoSampler:
# Initialize TeaCache if enabled # Initialize TeaCache if enabled
if teacache_args is not None: if teacache_args is not None:
# Check if dimensions have changed since last run
if (not hasattr(transformer, 'last_dimensions') or
transformer.last_dimensions != (height, width, num_frames) or
not hasattr(transformer, 'last_frame_count') or
transformer.last_frame_count != num_frames):
# Reset TeaCache state on dimension change
transformer.cnt = 0
transformer.teacache_skipped_steps = 0
transformer.accumulated_rel_l1_distance = 0
transformer.previous_modulated_input = None
transformer.previous_residual = None
transformer.last_dimensions = (height, width, num_frames)
transformer.last_frame_count = num_frames
transformer.teacache_device = device
transformer.enable_teacache = True transformer.enable_teacache = True
transformer.cnt = 0
transformer.accumulated_rel_l1_distance = 0
transformer.teacache_skipped_steps_cond = transformer.teacache_skipped_steps_uncond =0
transformer.previous_modulated_input_cond = transformer.previous_modulated_input_uncond = None
transformer.previous_residual_cond = transformer.previous_residual_uncond = None
transformer.accumulated_rel_l1_distance_cond = transformer.accumulated_rel_l1_distance_uncond = 0
transformer.teacache_device = device
transformer.num_steps = steps transformer.num_steps = steps
transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"] transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"]
transformer.teacache_start_step = teacache_args["start_step"]
teacache_end_step = teacache_args["end_step"]
if teacache_end_step < 0:
teacache_end_step = steps - 1
transformer.teacache_end_step = teacache_end_step
else: else:
transformer.enable_teacache = False transformer.enable_teacache = False
@@ -1472,6 +1457,24 @@ class HyVideoSampler:
if denoise_strength < 1.0: if denoise_strength < 1.0:
input_latents *= VAE_SCALING_FACTOR input_latents *= VAE_SCALING_FACTOR
mask_latents = None
if mask is not None:
from einops import rearrange
target_video_length = mask.shape[0]
target_height = mask.shape[1]
target_width = mask.shape[2]
mask_length = (target_video_length - 1) // 4 + 1
mask_height = target_height // 8
mask_width = target_width // 8
mask = mask.unsqueeze(-1).unsqueeze(0)
mask = rearrange(mask, "b t h w c -> b c t h w")
print("mask shape", mask.shape)
mask_latents = torch.nn.functional.interpolate(mask, size=(mask_length, mask_height, mask_width))
mask_latents = mask_latents.to(device)
out_latents = model["pipe"]( out_latents = model["pipe"](
num_inference_steps=steps, num_inference_steps=steps,
height = target_height, height = target_height,
@@ -1481,8 +1484,12 @@ class HyVideoSampler:
cfg_start_percent=cfg_start_percent, cfg_start_percent=cfg_start_percent,
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,
fresca_args=fresca_args,
slg_args=slg_args,
embedded_guidance_scale=embedded_guidance_scale, embedded_guidance_scale=embedded_guidance_scale,
latents=input_latents, latents=input_latents,
mask_latents=mask_latents,
denoise_strength=denoise_strength, denoise_strength=denoise_strength,
prompt_embed_dict=hyvid_embeds, prompt_embed_dict=hyvid_embeds,
generator=generator, generator=generator,
@@ -1495,9 +1502,10 @@ class HyVideoSampler:
feta_args=feta_args, feta_args=feta_args,
leapfusion_img2vid = leapfusion_img2vid, leapfusion_img2vid = leapfusion_img2vid,
image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None, image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None,
neg_image_cond_latents = neg_image_cond_latents["samples"] * VAE_SCALING_FACTOR if neg_image_cond_latents is not None else None,
riflex_freq_index = riflex_freq_index, riflex_freq_index = riflex_freq_index,
i2v_stability = i2v_stability, i2v_stability = i2v_stability,
taylorseer = taylorseer_args, loop_args = loop_args,
) )
print_memory(device) print_memory(device)
@@ -1507,8 +1515,10 @@ class HyVideoSampler:
pass pass
if teacache_args is not None: if teacache_args is not None:
log.info(f"TeaCache skipped {transformer.teacache_skipped_steps} steps")
transformer.teacache_skipped_steps = 0 log.info(f"TeaCache skipped {transformer.teacache_skipped_steps_cond} cond steps")
if transformer.teacache_skipped_steps_uncond > 0:
log.info(f"TeaCache skipped {transformer.teacache_skipped_steps_uncond} uncond steps")
if force_offload: if force_offload:
if model["manual_offloading"]: if model["manual_offloading"]:
@@ -1915,7 +1925,9 @@ NODE_CLASS_MAPPINGS = {
"HyVideoI2VEncode": HyVideoI2VEncode, "HyVideoI2VEncode": HyVideoI2VEncode,
"HyVideoEncodeKeyframes": HyVideoEncodeKeyframes, "HyVideoEncodeKeyframes": HyVideoEncodeKeyframes,
"HyVideoTextEmbedBridge": HyVideoTextEmbedBridge, "HyVideoTextEmbedBridge": HyVideoTextEmbedBridge,
"HyVideoTaylorSeerOptions": HyVideoTaylorSeerOptions, "HyVideoLoopArgs": HyVideoLoopArgs,
"HunyuanVideoFresca": HunyuanVideoFresca,
"HunyuanVideoSLG": HunyuanVideoSLG
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoSampler": "HunyuanVideo Sampler", "HyVideoSampler": "HunyuanVideo Sampler",
@@ -1943,5 +1955,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoI2VEncode": "HyVideo I2V Encode", "HyVideoI2VEncode": "HyVideo I2V Encode",
"HyVideoEncodeKeyframes": "HyVideo Encode Keyframes", "HyVideoEncodeKeyframes": "HyVideo Encode Keyframes",
"HyVideoTextEmbedBridge": "HyVideo TextEmbed Bridge", "HyVideoTextEmbedBridge": "HyVideo TextEmbed Bridge",
"HyVideoTaylorSeerOptions": "HyVideo TaylorSeer Options", "HyVideoLoopArgs": "HyVideo Loop Args",
"HunyuanVideoFresca": "HunyuanVideo Fresca",
"HunyuanVideoSLG": "HunyuanVideo SLG",
} }
+66
View File
@@ -24,3 +24,69 @@ def print_memory(device):
log.info(f"-------------------------------") log.info(f"-------------------------------")
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False) #memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
#log.info(f"Memory Summary:\n{memory_summary}") #log.info(f"Memory Summary:\n{memory_summary}")
def optimized_scale(positive_flat, negative_flat):
# Calculate dot production
dot_product = torch.sum(positive_flat * negative_flat, dim=1, keepdim=True)
# Squared norm of uncondition
squared_norm = torch.sum(negative_flat ** 2, dim=1, keepdim=True) + 1e-8
# st_star = v_cond^T * v_uncond / ||v_uncond||^2
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
B, C, T, H, W = x_freq.shape
crow, ccol = H // 2, W // 2
# Initialize mask with high-frequency scaling factor
mask = torch.ones((B, C, T, 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