Add TeaCache

This commit is contained in:
kijai
2025-01-06 06:26:28 +02:00
parent 46e31f1a44
commit 6c70beda52
2 changed files with 178 additions and 94 deletions
+132 -93
View File
@@ -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,
# },
# }
+46 -1
View File
@@ -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",
}