Fixed: When there are multiple TeaCache nodes in a workflow, the rel_l1_thresh value is always the value of the last TeaCache node

This commit is contained in:
TangYanxin
2025-01-15 09:22:28 +08:00
parent a0750c4e6f
commit be41e9ad9f
+13 -5
View File
@@ -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"