diff --git a/nodes.py b/nodes.py index a10ce7d..c012afa 100644 --- a/nodes.py +++ b/nodes.py @@ -32,6 +32,8 @@ def teacache_flux_forward( attn_mask: Tensor = None, ) -> Tensor: patches_replace = transformer_options.get("patches_replace", {}) + rel_l1_thresh = transformer_options.get("teacache_rel_l1_thresh", {}) + if img.ndim != 3 or txt.ndim != 3: raise ValueError("Input img and txt tensors must have 3 dimensions.") @@ -66,7 +68,7 @@ def teacache_flux_forward( try: coefficients = [4.98651651e+02, -2.83781631e+02, 5.58554382e+01, -3.82021401e+00, 2.64230861e-01] self.accumulated_rel_l1_distance += poly1d(coefficients, ((modulated_inp-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean())) - if self.accumulated_rel_l1_distance < self.rel_l1_thresh: + if self.accumulated_rel_l1_distance < rel_l1_thresh: should_calc = False else: should_calc = True @@ -185,6 +187,7 @@ def teacache_hunyuanvideo_forward( transformer_options={}, ) -> Tensor: patches_replace = transformer_options.get("patches_replace", {}) + rel_l1_thresh = transformer_options.get("teacache_rel_l1_thresh", {}) initial_shape = list(img.shape) # running on sequences img @@ -230,7 +233,7 @@ def teacache_hunyuanvideo_forward( try: coefficients = [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02] self.accumulated_rel_l1_distance += poly1d(coefficients, ((modulated_inp-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean())) - if self.accumulated_rel_l1_distance < self.rel_l1_thresh: + if self.accumulated_rel_l1_distance < rel_l1_thresh: should_calc = False else: should_calc = True @@ -312,6 +315,7 @@ def teacache_ltxvmodel_forward( **kwargs ): patches_replace = transformer_options.get("patches_replace", {}) + rel_l1_thresh = transformer_options.get("teacache_rel_l1_thresh", {}) indices_grid = self.patchifier.get_grid( orig_num_frames=x.shape[2], @@ -395,7 +399,7 @@ def teacache_ltxvmodel_forward( try: coefficients = [2.14700694e+01, -1.28016453e+01, 2.31279151e+00, 7.92487521e-01, 9.69274326e-03] self.accumulated_rel_l1_distance += poly1d(coefficients, ((modulated_inp-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean())) - if self.accumulated_rel_l1_distance < self.rel_l1_thresh: + if self.accumulated_rel_l1_distance < rel_l1_thresh: should_calc = False else: should_calc = True @@ -473,8 +477,10 @@ class TeaCacheForImgGen: def apply_teacache(self, model, model_type: str, rel_l1_thresh: float): new_model = model.clone() + if 'transformer_options' not in new_model.model_options: + new_model.model_options['transformer_options'] = {} + new_model.model_options["transformer_options"]["teacache_rel_l1_thresh"] = rel_l1_thresh diffusion_model = new_model.get_model_object("diffusion_model") - diffusion_model.__class__.rel_l1_thresh = rel_l1_thresh if model_type == "flux": forward_name = "forward_orig" @@ -514,8 +520,10 @@ class TeaCacheForVidGen: def apply_teacache(self, model, model_type: str, rel_l1_thresh: float): new_model = model.clone() + if 'transformer_options' not in new_model.model_options: + new_model.model_options['transformer_options'] = {} + new_model.model_options["transformer_options"]["teacache_rel_l1_thresh"] = rel_l1_thresh diffusion_model = new_model.get_model_object("diffusion_model") - diffusion_model.__class__.rel_l1_thresh = rel_l1_thresh if model_type == "hunyuan_video": forward_name = "forward_orig"