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:
@@ -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
@@ -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)
|
||||
Reference in New Issue
Block a user