Add TeaCache
This commit is contained in:
+132
-93
@@ -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,
|
||||
# },
|
||||
# }
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user