Experimental TeaCache

I haven't figured out the whole coefficiency calculation part, however I noticed it's already super close to the time embed as it is, if we just skip the initial steps from the calculations. Seems to work okayish with the 1.3B model at least.
This commit is contained in:
kijai
2025-03-02 01:34:55 +02:00
parent df6d0b677f
commit fbb432ce0f
2 changed files with 202 additions and 65 deletions
+54 -27
View File
@@ -74,26 +74,29 @@ class WanVideoBlockSwap:
# return (kwargs, )
# class WanVideoTeaCache:
# @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 = "WanVideoWrapper"
# DESCRIPTION = "TeaCache settings for WanVideo to speed up inference"
class WanVideoTeaCache:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"rel_l1_thresh": ("FLOAT", {"default": 0.04, "min": 0.0, "max": 1.0, "step": 0.001,
"tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts"}),
"start_step": ("INT", {"default": 6, "min": 0, "max": 9999, "step": 1, "tooltip": "Start percentage of the steps to apply TeaCache"}),
},
}
RETURN_TYPES = ("TEACACHEARGS",)
RETURN_NAMES = ("teacache_args",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "WORK IN PROGRESS! Naive approach, currently does NOT use calculated coefficiencies. Speeds up inference by skipping steps based on input/output difference"
EXPERIMENTAL = True
# def process(self, rel_l1_thresh):
# teacache_args = {
# "rel_l1_thresh": rel_l1_thresh,
# }
# return (teacache_args,)
def process(self, rel_l1_thresh, start_step):
teacache_args = {
"rel_l1_thresh": rel_l1_thresh,
"start_step": start_step,
}
return (teacache_args,)
class WanVideoModel(comfy.model_base.BaseModel):
@@ -912,6 +915,7 @@ class WanVideoSampler:
"denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"feta_args": ("FETAARGS", ),
"context_options": ("WANVIDCONTEXT", ),
"teacache_args": ("TEACACHEARGS", ),
}
}
@@ -921,7 +925,7 @@ class WanVideoSampler:
CATEGORY = "WanVideoWrapper"
def process(self, model, text_embeds, image_embeds, shift, steps, cfg, seed, scheduler, riflex_freq_index,
force_offload=True, samples=None, feta_args=None, denoise_strength=1.0, context_options=None):
force_offload=True, samples=None, feta_args=None, denoise_strength=1.0, context_options=None, teacache_args=None):
patcher = model
model = model.model
transformer = model.diffusion_model
@@ -1102,7 +1106,15 @@ class WanVideoSampler:
set_num_frames(latent_video_length)
enable_enhance()
else:
disable_enhance()
disable_enhance()
# Initialize TeaCache if enabled
if teacache_args is not None:
transformer.enable_teacache = True
transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"]
transformer.teacache_start_step = teacache_args["start_step"]
else:
transformer.enable_teacache = False
mm.soft_empty_cache()
gc.collect()
@@ -1163,10 +1175,10 @@ class WanVideoSampler:
partial_latent_model_input = [latent_model_input[0][:, c, :, :]]
# Model inference - returns [frames, channels, height, width]
noise_pred_cond = transformer(
partial_latent_model_input, t=timestep, **arg_c)[0].to(intermediate_device)
partial_latent_model_input, t=timestep, current_step=i,**arg_c)[0].to(intermediate_device)
if cfg[i] != 1.0:
noise_pred_uncond = transformer(
partial_latent_model_input, t=timestep, **arg_null)[0].to(intermediate_device)
partial_latent_model_input, t=timestep, current_step=i,**arg_null)[0].to(intermediate_device)
noise_pred_context = noise_pred_uncond + cfg[i] * (
noise_pred_cond - noise_pred_uncond)
@@ -1193,10 +1205,20 @@ class WanVideoSampler:
else:
#model inference start
noise_pred_cond = transformer(
latent_model_input, t=timestep, **arg_c)[0].to(intermediate_device)
latent_model_input,
t=timestep,
current_step=i,
is_uncond=False,
**arg_c
)[0].to(intermediate_device)
if cfg[i] != 1.0:
noise_pred_uncond = transformer(
latent_model_input, t=timestep, **arg_null)[0].to(intermediate_device)
latent_model_input,
t=timestep,
current_step=i,
is_uncond=True,
**arg_null
)[0].to(intermediate_device)
noise_pred = noise_pred_uncond + cfg[i] * (
noise_pred_cond - noise_pred_uncond)
@@ -1223,6 +1245,9 @@ class WanVideoSampler:
pbar.update(1)
del latent_model_input, timestep
if teacache_args is not None:
log.info(f"TeaCache skipped: {transformer.teacache_skipped_cond_steps} cond steps, {transformer.teacache_skipped_uncond_steps} uncond steps")
if transformer.attention_mode == "spargeattn_tune":
saved_state_dict = extract_sparse_attention_state_dict(transformer)
torch.save(saved_state_dict, "sparge_wan.pt")
@@ -1431,7 +1456,8 @@ NODE_CLASS_MAPPINGS = {
"WanVideoLoraSelect": WanVideoLoraSelect,
"WanVideoLoraBlockEdit": WanVideoLoraBlockEdit,
"WanVideoEnhanceAVideo": WanVideoEnhanceAVideo,
"WanVideoContextOptions": WanVideoContextOptions
"WanVideoContextOptions": WanVideoContextOptions,
"WanVideoTeaCache": WanVideoTeaCache
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -1452,5 +1478,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoLoraSelect": "WanVideo Lora Select",
"WanVideoLoraBlockEdit": "WanVideo Lora Block Edit",
"WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video",
"WanVideoContextOptions": "WanVideo Context Options"
"WanVideoContextOptions": "WanVideo Context Options",
"WanVideoTeaCache": "WanVideo TeaCache"
}
+148 -38
View File
@@ -10,11 +10,19 @@ from ...enhance_a_video.enhance import get_feta_scores
from ...enhance_a_video.globals import is_enhance_enabled
from .attention import attention
import numpy as np
__all__ = ['WanModel']
from tqdm import tqdm
from ...utils import log
def poly1d(coefficients, x):
result = torch.zeros_like(x)
for i, coeff in enumerate(coefficients):
result += coeff * (x ** (len(coefficients) - 1 - i))
return result.abs()
def sinusoidal_embedding_1d(dim, position):
# preprocess
assert dim % 2 == 0
@@ -491,6 +499,15 @@ class WanModel(ModelMixin, ConfigMixin):
self.offload_txt_emb = False
self.offload_img_emb = False
#init TeaCache variables
self.enable_teacache = False
self.teacache_counter = 0
self.rel_l1_thresh = 0.15
self.teacache_start_step= 2
# self.l1_history_x = []
# self.l1_history_temb = []
# self.l1_history_rescaled = []
# embeddings
self.patch_embedding = nn.Conv3d(
in_dim, dim, kernel_size=patch_size, stride=patch_size)
@@ -546,6 +563,8 @@ class WanModel(ModelMixin, ConfigMixin):
y=None,
device=torch.device('cuda'),
freqs=None,
current_step=0,
is_uncond=False
):
r"""
Forward pass through the diffusion model
@@ -567,7 +586,7 @@ class WanModel(ModelMixin, ConfigMixin):
Returns:
List[Tensor]:
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
"""
"""
if self.model_type == 'i2v':
assert clip_fea is not None and y is not None
# params
@@ -618,22 +637,107 @@ class WanModel(ModelMixin, ConfigMixin):
if self.offload_img_emb:
self.img_emb.to(self.offload_device, non_blocking=True)
# arguments
kwargs = dict(
e=e0,
seq_lens=seq_lens,
grid_sizes=grid_sizes,
freqs=freqs,
context=context,
context_lens=context_lens)
should_calc = True
if self.enable_teacache and current_step >= self.teacache_start_step:
if current_step == self.teacache_start_step:
log.info("TeaCache: Initializing TeaCache variables")
should_calc = True
self.accumulated_rel_l1_distance_cond = 0
self.accumulated_rel_l1_distance_uncond = 0
self.teacache_skipped_cond_steps = 0
self.teacache_skipped_uncond_steps = 0
else:
#coefficients = [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02] # Hunyuan
#coefficients = [-3.10658903e+01, 2.54732368e+01, -5.92380459e+00, 1.75769064e+00, -3.61568434e-03] #Cog2b
#coefficients = [-1.53880483e+03, 8.43202495e+02, -1.34363087e+02, 7.97131516e+00, -5.23162339e-02] #Cog5b
#self.accumulated_rel_l1_distance += poly1d(coefficients, ((e0-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean()))
prev_input = self.previous_modulated_input_uncond if is_uncond else self.previous_modulated_input_cond
acc_distance_attr = 'accumulated_rel_l1_distance_uncond' if is_uncond else 'accumulated_rel_l1_distance_cond'
for b, block in enumerate(self.blocks):
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
block.to(self.main_device)
x = block(x, **kwargs)
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
block.to(self.offload_device, non_blocking=True)
temb_relative_l1 = relative_l1_distance(prev_input, e0)
setattr(self, acc_distance_attr, getattr(self, acc_distance_attr) + temb_relative_l1)
if getattr(self, acc_distance_attr) < self.rel_l1_thresh:
should_calc = False
self.teacache_counter += 1
else:
should_calc = True
setattr(self, acc_distance_attr, 0)
# if current_step > 0:
# temb_relative_l1 = relative_l1_distance(self.previous_modulated_input, e0)
# print("temb_relative_l1 ", temb_relative_l1)
# self.l1_history_temb.append(temb_relative_l1.cpu())
if is_uncond:
self.previous_modulated_input_uncond = e0.clone()
if not should_calc:
x += self.previous_residual_uncond
#log.info(f"TeaCache: Skipping uncond step {current_step+1}")
self.teacache_skipped_cond_steps += 1
else:
self.previous_modulated_input_cond = e0.clone()
if not should_calc:
x += self.previous_residual_cond
#log.info(f"TeaCache: Skipping cond step {current_step+1}")
self.teacache_skipped_uncond_steps += 1
if not self.enable_teacache or (self.enable_teacache and should_calc):
if self.enable_teacache:
ori_hidden_states = x.clone()
# arguments
kwargs = dict(
e=e0,
seq_lens=seq_lens,
grid_sizes=grid_sizes,
freqs=freqs,
context=context,
context_lens=context_lens)
for b, block in enumerate(self.blocks):
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
block.to(self.main_device)
x = block(x, **kwargs)
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
block.to(self.offload_device, non_blocking=True)
if self.enable_teacache:
if is_uncond:
self.previous_residual_uncond = x - ori_hidden_states
else:
self.previous_residual_cond = x - ori_hidden_states
# if current_step > 0:
# import matplotlib.pyplot as plt
# x_relative_l1 = relative_l1_distance(x,ori_hidden_states)
# print("x_relative_l1 ", x_relative_l1)
# self.l1_history_x.append(x_relative_l1.cpu())
# # Rescale using polynomial fitting
# if len(self.l1_history_x) > 1:
# print("self.l1_history_temb ", self.l1_history_temb)
# rescaled_diffs = rescale_differences(
# self.l1_history_temb,
# self.l1_history_x
# )
# self.l1_history_rescaled = rescaled_diffs#.tolist()
# #print("x_relative_l1 ", x_relative_l1)
# if current_step == self.num_steps-1:
# plt.figure(figsize=(10,5))
# norm_x = normalize_values([x.item() for x in self.l1_history_x])
# norm_temb = normalize_values([x.item() for x in self.l1_history_temb])
# norm_rescaled = normalize_values(self.l1_history_rescaled)
# plt.plot(norm_x, label='Hidden States L1')
# plt.plot(norm_temb, label='Original Temb L1')
# plt.plot(norm_rescaled, label='Rescaled Temb L1')
# plt.title('Relative L1 Distances Over Time')
# plt.xlabel('Step')
# plt.ylabel('Normalized L1 Distance')
# plt.grid(True)
# plt.legend()
# plt.savefig('l1_distances_plot.png')
# plt.close()
# head
x = self.head(x, e)
@@ -667,26 +771,32 @@ class WanModel(ModelMixin, ConfigMixin):
out.append(u)
return out
# def init_weights(self):
# r"""
# Initialize model parameters using Xavier initialization.
# """
def relative_l1_distance(last_tensor, current_tensor):
l1_distance = torch.abs(last_tensor - current_tensor).mean()
norm = torch.abs(last_tensor).mean()
relative_l1_distance = l1_distance / norm
return relative_l1_distance.to(torch.float32)
# # basic init
# for m in self.modules():
# if isinstance(m, nn.Linear):
# nn.init.xavier_uniform_(m.weight)
# if m.bias is not None:
# nn.init.zeros_(m.bias)
def normalize_values(values):
min_val = min(values)
max_val = max(values)
if max_val == min_val:
return [0.0] * len(values)
return [(x - min_val) / (max_val - min_val) for x in values]
# # init embeddings
# nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
# for m in self.text_embedding.modules():
# if isinstance(m, nn.Linear):
# nn.init.normal_(m.weight, std=.02)
# for m in self.time_embedding.modules():
# if isinstance(m, nn.Linear):
# nn.init.normal_(m.weight, std=.02)
# # init output layer
# nn.init.zeros_(self.head.head.weight)
def rescale_differences(input_diffs, output_diffs):
"""Polynomial fitting between input and output differences"""
poly_degree = 4
if len(input_diffs) < 2:
return input_diffs
x = np.array([x.item() for x in input_diffs])
y = np.array([y.item() for y in output_diffs])
print("x ", x)
print("y ", y)
# Fit polynomial
coeffs = np.polyfit(x, y, poly_degree)
# Apply polynomial transformation
return np.polyval(coeffs, x)