Update TeaCache_Lumina2.py

This commit is contained in:
spawner
2025-06-07 14:11:20 +08:00
committed by GitHub
parent 88b9a499fb
commit aaecb467e2
+155 -56
View File
@@ -2,12 +2,15 @@ import torch
import numpy as np
from comfy.ldm.common_dit import pad_to_patch_size # noqa
from unittest.mock import patch
import re
DEFAULT_COEFFICIENTS = [393.76566581, -603.50993606, 209.10239044, -23.00726601, 0.86377344]
# referenced from https://github.com/spawner1145/TeaCache/blob/main/TeaCache4Lumina2/teacache_lumina2.py
# firstly transplanted by @fexli https://github.com/fexli
# retransplanted by @spawner1145 https://github.com/spawner1145
def teacache_forward_working(
self, x, timesteps, context, num_tokens, attention_mask=None, transformer_options={}, **kwargs
self, x, timesteps, context, num_tokens, attention_mask=None, transformer_options={}, **kwargs
):
if not hasattr(self, 'teacache_state'):
self.teacache_state = {
@@ -18,7 +21,7 @@ def teacache_forward_working(
}
if not isinstance(self.teacache_state.get("cache"), dict):
self.teacache_state["cache"] = {}
if self.teacache_state.get("num_steps") is None and transformer_options.get("num_steps") is not None:
self.teacache_state["num_steps"] = transformer_options.get("num_steps")
@@ -41,7 +44,8 @@ def teacache_forward_working(
enable_teacache = transformer_options.get('enable_teacache', False)
current_cache = None
modulated_inp = None
if enable_teacache:
cache_key = max_seq_len
if cache_key not in self.teacache_state['cache']:
@@ -51,41 +55,77 @@ def teacache_forward_working(
"previous_residual": None,
}
current_cache = self.teacache_state['cache'][cache_key]
modulated_inp = self.layers[0].adaLN_modulation(adaln_input.clone())[0]
try:
if self.layers and hasattr(self.layers[0], 'adaLN_modulation'):
mod_result = self.layers[0].adaLN_modulation(adaln_input.clone())
if isinstance(mod_result, (list, tuple)) and len(mod_result) > 0:
modulated_inp = mod_result[0]
elif torch.is_tensor(mod_result):
modulated_inp = mod_result
else:
raise ValueError("adaLN_modulation returned unexpected type or empty list/tuple")
else:
raise AttributeError("Layer 0 or adaLN_modulation not found")
except Exception as e:
print(f"Warning: TeaCache - Failed to get modulated_inp: {e}. Disabling cache for this step.")
enable_teacache = False
should_calc = True
modulated_inp = None
if current_cache:
current_cache["previous_modulated_input"] = None
current_cache["accumulated_rel_l1_distance"] = 0.0
if enable_teacache and modulated_inp is not None and current_cache is not None:
num_steps_in_state = self.teacache_state.get("num_steps")
if num_steps_in_state is None or num_steps_in_state == 0:
should_calc = True
if current_cache: current_cache["accumulated_rel_l1_distance"] = 0.0
current_cache["accumulated_rel_l1_distance"] = 0.0
elif self.teacache_state['cnt'] == 0 or self.teacache_state['cnt'] == num_steps_in_state - 1:
should_calc = True
if current_cache: current_cache["accumulated_rel_l1_distance"] = 0.0
current_cache["accumulated_rel_l1_distance"] = 0.0
else:
if current_cache and current_cache.get("previous_modulated_input") is not None:
coefficients = [393.76566581, -603.50993606, 209.10239044, -23.00726601,
0.86377344]
rescale_func = np.poly1d(coefficients)
if current_cache.get("previous_modulated_input") is not None:
# coefficients = [393.76566581, -603.50993606, 209.10239044, -23.00726601, 0.86377344]
coefficients = transformer_options.get('coefficients', DEFAULT_COEFFICIENTS)
if not isinstance(coefficients, (list, tuple)):
coefficients = DEFAULT_COEFFICIENTS
try:
rescale_func = np.poly1d(coefficients)
except Exception as e:
print(f"Warning: TeaCache np.poly1d failed with coefficients {coefficients}: {e}. Using default for this step.")
rescale_func = np.poly1d(DEFAULT_COEFFICIENTS)
prev_mod_input = current_cache["previous_modulated_input"]
prev_mean = prev_mod_input.abs().mean()
if prev_mean.item() > 1e-9:
rel_l1_change = ((modulated_inp - prev_mod_input).abs().mean() / prev_mean).cpu().item()
if prev_mod_input.shape != modulated_inp.shape:
print(f"Warning: TeaCache - modulated input shape mismatch: prev={prev_mod_input.shape}, curr={modulated_inp.shape}. Forcing recalculation.")
should_calc = True
current_cache["accumulated_rel_l1_distance"] = 0.0
rel_l1_change = float('inf')
else:
rel_l1_change = 0.0 if modulated_inp.abs().mean().item() < 1e-9 else float('inf')
current_cache["accumulated_rel_l1_distance"] += rescale_func(rel_l1_change)
prev_mean = prev_mod_input.abs().mean()
if prev_mean.item() > 1e-9:
rel_l1_change = ((modulated_inp - prev_mod_input).abs().mean() / prev_mean).cpu().item()
else:
rel_l1_change = 0.0 if modulated_inp.abs().mean().item() < 1e-9 else float('inf')
rescaled_value = rescale_func(rel_l1_change)
if np.isnan(rescaled_value) or np.isinf(rescaled_value):
current_cache["accumulated_rel_l1_distance"] = float('inf')
else:
current_cache["accumulated_rel_l1_distance"] += rescaled_value
if current_cache["accumulated_rel_l1_distance"] < transformer_options.get('rel_l1_thresh', 0.3):
should_calc = False
else:
should_calc = True
current_cache["accumulated_rel_l1_distance"] = 0.0
if current_cache["accumulated_rel_l1_distance"] < transformer_options.get('rel_l1_thresh', 0.3):
should_calc = False
else:
should_calc = True
current_cache["accumulated_rel_l1_distance"] = 0.0
else:
should_calc = True
if current_cache: current_cache["accumulated_rel_l1_distance"] = 0.0
current_cache["accumulated_rel_l1_distance"] = 0.0
if current_cache:
current_cache["previous_modulated_input"] = modulated_inp.clone()
current_cache["previous_modulated_input"] = modulated_inp.clone()
if self.teacache_state.get('uncond_seq_len') is None:
self.teacache_state['uncond_seq_len'] = cache_key
@@ -95,9 +135,20 @@ def teacache_forward_working(
if self.teacache_state['cnt'] >= num_steps_in_state:
self.teacache_state['cnt'] = 0
if enable_teacache and not should_calc and current_cache and current_cache.get("previous_residual") is not None:
processed_x = x + current_cache["previous_residual"]
can_reuse_residual = (enable_teacache and
not should_calc and
current_cache and
current_cache.get("previous_residual") is not None and
current_cache["previous_residual"].shape == x.shape) # 检查形状
if can_reuse_residual:
processed_x = x + current_cache["previous_residual"]
else:
if enable_teacache and not should_calc and current_cache and current_cache.get("previous_residual") is not None and current_cache["previous_residual"].shape != x.shape:
print(f"Warning: TeaCache - Residual shape mismatch: cache={current_cache['previous_residual'].shape}, input={x.shape}. Forcing recalculation.")
if current_cache:
current_cache["accumulated_rel_l1_distance"] = 0.0
original_x = x.clone()
current_x_for_processing = x
for layer in self.layers:
@@ -107,7 +158,7 @@ def teacache_forward_working(
current_cache["previous_residual"] = current_x_for_processing - original_x
current_cache["accumulated_rel_l1_distance"] = 0.0
processed_x = current_x_for_processing
output = self.final_layer(processed_x, adaln_input)
output = self.unpatchify(output, img_size, cap_size, return_tensor=True)[:, :, :h_img, :w_img]
@@ -123,9 +174,14 @@ class TeaCache_Lumina2:
"model": ("MODEL",),
"rel_l1_thresh": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 10.0, "step": 0.001}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01,
"tooltip": "The start percentage of the steps that will apply TeaCache."}),
"tooltip": "The start percentage of the steps that will apply TeaCache. / TeaCache开始应用的步数百分比。"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01,
"tooltip": "The end percentage of the steps that will apply TeaCache."})
"tooltip": "The end percentage of the steps that will apply TeaCache. / TeaCache停止应用的步数百分比。"}),
"coefficients_string": ("STRING", {
"multiline": True,
"default": str(DEFAULT_COEFFICIENTS),
"tooltip": "Coefficients for np.poly1d. Format: 393.7, -603.5, 209.1, -23.0, 0.86 (with or without brackets []) / 用于 np.poly1d 的系数。格式: 393.7, -603.5, 209.1, -23.0, 0.86 (可带或不带方括号[])"
}),
}
}
@@ -133,20 +189,43 @@ class TeaCache_Lumina2:
FUNCTION = "patch_teacache"
CATEGORY = "utils"
def patch_teacache(self, model, rel_l1_thresh, start_percent, end_percent):
# start_percent = 0.0
# end_percent = 1.0
def patch_teacache(self, model, rel_l1_thresh, start_percent, end_percent, coefficients_string):
if rel_l1_thresh == 0:
try:
diffusion_model = model.get_model_object("diffusion_model")
if hasattr(diffusion_model, 'teacache_state'):
delattr(diffusion_model, 'teacache_state')
except:
pass
return (model,)
parsed_coefficients = DEFAULT_COEFFICIENTS
try:
s = re.sub(r'[\[\]\s]', '', coefficients_string.strip())
if not s:
parsed_coefficients = DEFAULT_COEFFICIENTS
else:
coeff_list = [float(item) for item in s.split(',') if item]
if not coeff_list:
parsed_coefficients = DEFAULT_COEFFICIENTS
else:
parsed_coefficients = coeff_list
except ValueError:
print(f"Warning: TeaCache - Could not parse coefficients string: '{coefficients_string}'. Using default coefficients: {DEFAULT_COEFFICIENTS}")
parsed_coefficients = DEFAULT_COEFFICIENTS
except Exception as e:
print(f"Warning: TeaCache - Error parsing coefficients '{coefficients_string}': {e}. Using default coefficients: {DEFAULT_COEFFICIENTS}")
parsed_coefficients = DEFAULT_COEFFICIENTS
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"]["cache"] = {} # 初始化空缓存
new_model.model_options["transformer_options"]["cache"] = {}
new_model.model_options["transformer_options"]["uncond_seq_len"] = None
new_model.model_options["transformer_options"]["rel_l1_thresh"] = rel_l1_thresh
new_model.model_options["transformer_options"]["coefficients"] = parsed_coefficients
diffusion_model = new_model.get_model_object("diffusion_model")
@@ -178,32 +257,51 @@ class TeaCache_Lumina2:
print("warning: TeaCache - 'sample_sigmas' not found in c.transformer_options.TeaCache might not work correctly.")
c_condition_dict["transformer_options"]["enable_teacache"] = False
c_condition_dict["transformer_options"]["num_steps"] = 1
if hasattr(diffusion_model, 'teacache_state'):
delattr(diffusion_model, 'teacache_state')
else:
sigmas = c_condition_dict["transformer_options"]["sample_sigmas"]
total_sampler_steps = len(sigmas)
c_condition_dict["transformer_options"]["num_steps"] = total_sampler_steps
if hasattr(diffusion_model, 'teacache_state') and diffusion_model.teacache_state is not None:
if diffusion_model.teacache_state.get("num_steps") != total_sampler_steps:
diffusion_model.teacache_state['num_steps'] = total_sampler_steps
if total_sampler_steps > 0:
c_condition_dict["transformer_options"]["num_steps"] = total_sampler_steps
else:
c_condition_dict["transformer_options"]["num_steps"] = 1
c_condition_dict["transformer_options"]["enable_teacache"] = False
matched_step_index = (sigmas == timestep[0]).nonzero()
if len(matched_step_index) > 0:
current_step_index = matched_step_index.item()
if hasattr(diffusion_model, 'teacache_state') and diffusion_model.teacache_state is not None:
if diffusion_model.teacache_state.get("num_steps") != total_sampler_steps and total_sampler_steps > 0:
delattr(diffusion_model, 'teacache_state')
c_condition_dict["transformer_options"]["cache"] = {}
current_timestep = timestep[0].to(device=sigmas.device, dtype=sigmas.dtype)
# 使用 torch.isclose 和 any() 来处理浮点数比较
close_mask = torch.isclose(sigmas, current_timestep, atol=1e-6)
if close_mask.any():
matched_step_index = torch.nonzero(close_mask, as_tuple=True)[0]
current_step_index = matched_step_index[0].item()
else:
current_step_index = 0
if total_sampler_steps > 1:
for i in range(total_sampler_steps - 1):
if (sigmas[i] - timestep[0]) * (sigmas[i + 1] - timestep[0]) <= 0:
current_step_index = i
break
try:
indices = torch.where(sigmas >= current_timestep)[0]
if len(indices) > 0:
current_step_index = indices[-1].item()
else:
current_step_index = total_sampler_steps -1
except:
for i in range(total_sampler_steps - 1):
if (sigmas[i] - current_timestep) * (sigmas[i + 1] - current_timestep) <= 0:
current_step_index = i
break
current_percent = 0.0
if total_sampler_steps > 1:
current_percent = current_step_index / (total_sampler_steps - 1)
elif total_sampler_steps == 1:
current_percent = 0.0
else:
current_percent = 0.0
current_percent = current_step_index / max(1, (total_sampler_steps - 1))
# elif total_sampler_steps == 1:
# current_percent = 0.0
elif total_sampler_steps <= 0 :
# current_percent = 0.0
c_condition_dict["transformer_options"]["enable_teacache"] = False
if start_percent <= current_percent <= end_percent and total_sampler_steps > 0:
@@ -212,10 +310,11 @@ class TeaCache_Lumina2:
c_condition_dict["transformer_options"]["enable_teacache"] = False
if current_step_index == 0:
if (1 in cond_or_uncond) and hasattr(diffusion_model, 'teacache_state'):
delattr(diffusion_model, 'teacache_state')
elif (0 in cond_or_uncond) and hasattr(diffusion_model, 'teacache_state'):
delattr(diffusion_model, 'teacache_state')
# print("Debug: TeaCache - Resetting state at step 0")
if hasattr(diffusion_model, 'teacache_state'):
delattr(diffusion_model, 'teacache_state')
c_condition_dict["transformer_options"]["cache"] = {}
with context_patch_manager:
return model_function(input_val, timestep, **c_condition_dict)