From 6c70beda520da11439b905eceedc4e3690e328d2 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 6 Jan 2025 06:26:28 +0200 Subject: [PATCH] Add TeaCache --- hyvideo/modules/models.py | 225 ++++++++++++++++++++++---------------- nodes.py | 47 +++++++- 2 files changed, 178 insertions(+), 94 deletions(-) diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 11d1ad3..e6d5c3f 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -3,7 +3,8 @@ from einops import rearrange import torch import torch.nn as nn -import torch.nn.functional as F + +import numpy as np from diffusers.models import ModelMixin from diffusers.configuration_utils import ConfigMixin, register_to_config @@ -671,11 +672,23 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): get_activation_layer("silu"), **factory_kwargs, ) + #init block swap variables self.double_blocks_to_swap = -1 self.single_blocks_to_swap = -1 self.offload_txt_in = False self.offload_img_in = False + #init TeaCache variables + self.enable_teacache = False + self.cnt = 0 + self.num_steps = 0 + self.rel_l1_thresh = 0.15 + self.accumulated_rel_l1_distance = 0 + self.previous_modulated_input = None + self.previous_residual = None + self.last_dimensions = None + self.last_frame_count = None + # 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): print(f"Swapping {double_blocks_to_swap + 1} double blocks and {single_blocks_to_swap + 1} single blocks") @@ -866,6 +879,30 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): stg_block_idx: int = -1, return_dict: bool = True, ) -> Union[torch.Tensor, Dict[str, torch.Tensor]]: + + def _process_double_blocks(img, txt, vec, block_args): + for b, block in enumerate(self.double_blocks): + if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0: + block.to(self.main_device) + + img, txt = block(img, txt, vec, *block_args) + + if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0: + block.to(self.offload_device, non_blocking=True) + return img, txt + + 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): + if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0: + block.to(self.main_device) + + curr_stg_mode = stg_mode if b == stg_block_idx else None + x = block(x, vec, txt_seq_len, *block_args, curr_stg_mode) + + if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0: + block.to(self.offload_device, non_blocking=True) + return x + out = {} img = x txt = text_states @@ -877,6 +914,21 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): ) set_num_frames(img.shape[2]) + current_dims = (ot, oh, ow) + + # Check if dimensions changed since last run + if not hasattr(self, 'last_dims') or self.last_dims != current_dims: + # Reset TeaCache state on dimension change + self.cnt = 0 + self.accumulated_rel_l1_distance = 0 + self.previous_modulated_input = None + self.previous_residual = None + self.last_dims = current_dims + + out = {} + img = x + txt = text_states + # Prepare modulation vectors. vec = self.time_in(t) @@ -931,57 +983,70 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): cu_seqlens_kv = cu_seqlens_q freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None - # --------------------- Pass through DiT blocks ------------------------ - for b, block in enumerate(self.double_blocks): - if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0: - #print(f"Moving double_block {b} to main device") - block.to(self.main_device) - double_block_args = [ - img, - txt, - vec, - cu_seqlens_q, - cu_seqlens_kv, - max_seqlen_q, - max_seqlen_kv, - freqs_cis, - attn_mask - ] - img, txt = block(*double_block_args) - if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0: - #print(f"Moving double_block {b} to offload device") - block.to(self.offload_device, non_blocking=True) + block_args = [cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv, freqs_cis, attn_mask] - # Merge txt and img to pass through single stream blocks. - x = torch.cat((img, txt), 1) - if len(self.single_blocks) > 0: - for b, block in enumerate(self.single_blocks): - if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0: - #print(f"Moving single_block {b} to main device") - #mm.soft_empty_cache() - block.to(self.main_device) - curr_stg_mode = stg_mode if b == stg_block_idx else None - single_block_args = [ - x, - vec, - txt_seq_len, - cu_seqlens_q, - cu_seqlens_kv, - max_seqlen_q, - max_seqlen_kv, - (freqs_cos, freqs_sin), - attn_mask, - curr_stg_mode, - ] + #tea_cache + if self.enable_teacache: + inp = img.clone() + vec_ = vec.clone() + txt_ = txt.clone() + self.double_blocks[0].to(self.main_device) + ( + img_mod1_shift, + img_mod1_scale, + img_mod1_gate, + img_mod2_shift, + img_mod2_scale, + img_mod2_gate, + ) = self.double_blocks[0].img_mod(vec_).chunk(6, dim=-1) + normed_inp = self.double_blocks[0].img_norm1(inp) + modulated_inp = modulate( + normed_inp, shift=img_mod1_shift, scale=img_mod1_scale + ) - x = block(*single_block_args) - if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0: - #print(f"Moving single_block {b} to offload device") - #mm.soft_empty_cache() - block.to(self.offload_device, non_blocking=True) + if self.cnt == 0 or self.cnt == self.num_steps-1: + should_calc = True + self.accumulated_rel_l1_distance = 0 + self.previous_modulated_input = modulated_inp.clone() + else: + coefficients = [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02] + 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 self.accumulated_rel_l1_distance < self.rel_l1_thresh: + should_calc = False + else: + should_calc = True + self.accumulated_rel_l1_distance = 0 + self.previous_modulated_input = modulated_inp.clone() + self.cnt += 1 + if self.cnt == self.num_steps: + self.cnt = 0 - img = x[:, :img_seq_len, ...] + if not should_calc and self.previous_residual is not None: + # Verify tensor dimensions match before adding + if img.shape == self.previous_residual.shape: + img = img + self.previous_residual + else: + should_calc = True # Force recalculation if dimensions don't match + + if should_calc: + ori_img = img.clone() + # 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, ...] + self.previous_residual = img - ori_img + 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, ...] # ---------------------------- Final layer ------------------------------ img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels) @@ -1007,52 +1072,26 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): return imgs - def params_count(self): - counts = { - "double": sum( - [ - sum(p.numel() for p in block.img_attn_qkv.parameters()) - + sum(p.numel() for p in block.img_attn_proj.parameters()) - + sum(p.numel() for p in block.img_mlp.parameters()) - + sum(p.numel() for p in block.txt_attn_qkv.parameters()) - + sum(p.numel() for p in block.txt_attn_proj.parameters()) - + sum(p.numel() for p in block.txt_mlp.parameters()) - for block in self.double_blocks - ] - ), - "single": sum( - [ - sum(p.numel() for p in block.linear1.parameters()) - + sum(p.numel() for p in block.linear2.parameters()) - for block in self.single_blocks - ] - ), - "total": sum(p.numel() for p in self.parameters()), - } - counts["attn+mlp"] = counts["double"] + counts["single"] - return counts - - ################################################################################# # HunyuanVideo Configs # ################################################################################# -HUNYUAN_VIDEO_CONFIG = { - "HYVideo-T/2": { - "mm_double_blocks_depth": 20, - "mm_single_blocks_depth": 40, - "rope_dim_list": [16, 56, 56], - "hidden_size": 3072, - "heads_num": 24, - "mlp_width_ratio": 4, - }, - "HYVideo-T/2-cfgdistill": { - "mm_double_blocks_depth": 20, - "mm_single_blocks_depth": 40, - "rope_dim_list": [16, 56, 56], - "hidden_size": 3072, - "heads_num": 24, - "mlp_width_ratio": 4, - "guidance_embed": True, - }, -} +# HUNYUAN_VIDEO_CONFIG = { +# "HYVideo-T/2": { +# "mm_double_blocks_depth": 20, +# "mm_single_blocks_depth": 40, +# "rope_dim_list": [16, 56, 56], +# "hidden_size": 3072, +# "heads_num": 24, +# "mlp_width_ratio": 4, +# }, +# "HYVideo-T/2-cfgdistill": { +# "mm_double_blocks_depth": 20, +# "mm_single_blocks_depth": 40, +# "rope_dim_list": [16, 56, 56], +# "hidden_size": 3072, +# "heads_num": 24, +# "mlp_width_ratio": 4, +# "guidance_embed": True, +# }, +# } diff --git a/nodes.py b/nodes.py index 8c3fd1f..2f2d801 100644 --- a/nodes.py +++ b/nodes.py @@ -175,6 +175,27 @@ class HyVideoSTG: def setargs(self, **kwargs): return (kwargs, ) +class HyVideoTeaCache: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "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"}), + }, + } + RETURN_TYPES = ("TEACACHEARGS",) + RETURN_NAMES = ("teacache_args",) + FUNCTION = "process" + CATEGORY = "HunyuanVideoWrapper" + DESCRIPTION = "TeaCache settings for HunyuanVideo to speed up inference" + + def process(self, rel_l1_thresh): + teacache_args = { + "rel_l1_thresh": rel_l1_thresh, + } + return (teacache_args,) + class HyVideoModel(comfy.model_base.BaseModel): def __init__(self, *args, **kwargs): @@ -1058,6 +1079,7 @@ class HyVideoSampler: "stg_args": ("STGARGS", ), "context_options": ("COGCONTEXT", ), "feta_args": ("FETAARGS", ), + "teacache_args": ("TEACACHEARGS", ) } } @@ -1067,7 +1089,7 @@ class HyVideoSampler: CATEGORY = "HunyuanVideoWrapper" 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): model = model.model device = mm.get_torch_device() @@ -1129,6 +1151,27 @@ class HyVideoSampler: elif model["manual_offloading"]: transformer.to(device) + # Initialize TeaCache if enabled + 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.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.enable_teacache = True + transformer.num_steps = steps + transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"] + else: + transformer.enable_teacache = False + mm.soft_empty_cache() gc.collect() @@ -1408,6 +1451,7 @@ NODE_CLASS_MAPPINGS = { "HyVideoTextEmbedsLoad": HyVideoTextEmbedsLoad, "HyVideoContextOptions": HyVideoContextOptions, "HyVideoEnhanceAVideo": HyVideoEnhanceAVideo, + "HyVideoTeaCache": HyVideoTeaCache, } NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoSampler": "HunyuanVideo Sampler", @@ -1430,4 +1474,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoTextEmbedsLoad": "HunyuanVideo TextEmbeds Load", "HyVideoContextOptions": "HunyuanVideo Context Options", "HyVideoEnhanceAVideo": "HunyuanVideo Enhance A Video", + "HyVideoTeaCache": "HunyuanVideo TeaCache", }