Update TeaCache_Lumina2.py
This commit is contained in:
+155
-56
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user