From fbb432ce0f6e810420e4e5fcd92f716e98804d5a Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 2 Mar 2025 01:34:55 +0200 Subject: [PATCH] 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. --- nodes.py | 81 +++++++++++------ wanvideo/modules/model.py | 186 ++++++++++++++++++++++++++++++-------- 2 files changed, 202 insertions(+), 65 deletions(-) diff --git a/nodes.py b/nodes.py index e57ae95..36a5ca7 100644 --- a/nodes.py +++ b/nodes.py @@ -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" } diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index e30c245..79e062d 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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) \ No newline at end of file