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,
# },
# }