From c3735392faff956e65e7ee23433ad690075c5fe9 Mon Sep 17 00:00:00 2001 From: spawner Date: Sat, 7 Jun 2025 00:16:24 +0800 Subject: [PATCH] Create tc_test.py --- tc_test.py | 399 +++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 399 insertions(+) create mode 100644 tc_test.py diff --git a/tc_test.py b/tc_test.py new file mode 100644 index 0000000..e23b1ec --- /dev/null +++ b/tc_test.py @@ -0,0 +1,399 @@ +import torch +import numpy as np +import time +import json +import os +from unittest.mock import patch +import folder_paths + +try: # lpips分析图像差异距离 + import lpips + LPIPS_AVAILABLE = True +except ImportError: + LPIPS_AVAILABLE = False + +try: # 贝叶斯 + from skopt import Optimizer + from skopt.space import Real + SKOPT_AVAILABLE = True +except ImportError: + SKOPT_AVAILABLE = False + +# 全局状态管理 +class TeaCacheStateManager: + _instance = None + def __new__(cls): + if cls._instance is None: + cls._instance = super(TeaCacheStateManager, cls).__new__(cls) + cls._instance.runs = {} + cls._instance.lpips_model = None + cls._instance.baseline_image = None + return cls._instance + + def register_run(self, run_id, params): + self.runs[run_id] = {"start_time": time.time(), "params": params, "results": {}, "cache": {}, "cnt": 0} + + def get_run_data(self, run_id): + return self.runs.get(run_id) + + def cleanup_run(self, run_id): + if run_id in self.runs: + del self.runs[run_id] + + def set_lpips_model(self, model): + self.lpips_model = model + + def get_lpips_model(self): + return self.lpips_model + + def set_baseline_image(self, image_tensor): + self.baseline_image = image_tensor + + def get_baseline_image(self): + return self.baseline_image + +STATE_MANAGER = TeaCacheStateManager() + +# 缓存逻辑 +def teacache_forward_analysis(self, x, timesteps, context, num_tokens, **kwargs): + transformer_options = kwargs.get('transformer_options', {}) + run_id = transformer_options.get('teacache_run_id') + if not run_id: return self.forward_original(x, timesteps, context, num_tokens, **kwargs) + + run_data = STATE_MANAGER.get_run_data(run_id) + if not run_data: return self.forward_original(x, timesteps, context, num_tokens, **kwargs) + + params = run_data['params'] + + if run_data.get("num_steps") is None and transformer_options.get("num_steps") is not None: + run_data["num_steps"] = transformer_options.get("num_steps") + + cap_feats, cap_mask = context, kwargs.get('attention_mask') + bs, c_channels, h_img, w_img = x.shape + if hasattr(self, 'pad_to_patch_size'): x = self.pad_to_patch_size(x, (self.patch_size, self.patch_size)) + else: + ph = (self.patch_size - h_img % self.patch_size) % self.patch_size + pw = (self.patch_size - w_img % self.patch_size) % self.patch_size + x = torch.nn.functional.pad(x, (0, pw, 0, ph)) + t = (1.0 - timesteps).to(dtype=x.dtype) + t_emb = self.t_embedder(t, dtype=x.dtype) + adaln_input = t_emb + if cap_feats is not None: cap_feats = self.cap_embedder(cap_feats) + x, mask, img_size, cap_size, freqs_cis = self.patchify_and_embed(x, cap_feats, cap_mask, t_emb, num_tokens) + freqs_cis = freqs_cis.to(x.device) + max_seq_len = x.shape[1] + + should_calc = True + num_steps_in_state = run_data.get("num_steps") + cnt = run_data.get('cnt', 0) + + if num_steps_in_state is None or num_steps_in_state == 0 or cnt == 0 or cnt == num_steps_in_state - 1: + should_calc = True + else: + run_data.setdefault('cache', {}) + current_cache = run_data['cache'].setdefault(max_seq_len, {"accumulated_rel_l1_distance": 0.0, "previous_modulated_input": None, "previous_residual": None}) + if current_cache.get("previous_modulated_input") is not None: + modulated_inp = self.layers[0].adaLN_modulation(adaln_input.clone())[0] + coefficients = params.get("coefficients_to_use", []) + rescale_func = np.poly1d(coefficients) + prev_mod_input = current_cache["previous_modulated_input"] + prev_mean = prev_mod_input.abs().mean().item() + rel_l1_change = ((modulated_inp - prev_mod_input).abs().mean() / prev_mean).cpu().item() if prev_mean > 1e-9 else float('inf') + current_cache["accumulated_rel_l1_distance"] += rescale_func(rel_l1_change) + + if current_cache["accumulated_rel_l1_distance"] < params.get('rel_l1_thresh', 0.3): should_calc = False + else: + should_calc = True + current_cache["accumulated_rel_l1_distance"] = 0.0 + current_cache["previous_modulated_input"] = modulated_inp.clone() + else: should_calc = True + + if should_calc: + if max_seq_len != run_data.get('uncond_seq_len'): + run_data['results'].setdefault("total_inferences", 0) + run_data['results']["total_inferences"] += 1 + original_x = x.clone() + processed_x = x + for layer in self.layers: processed_x = layer(processed_x, mask, freqs_cis, adaln_input) + run_data.setdefault('cache', {}) + current_cache = run_data['cache'].setdefault(max_seq_len, {}) + current_cache["previous_residual"] = processed_x - original_x + current_cache["accumulated_rel_l1_distance"] = 0.0 + if current_cache.get("previous_modulated_input") is None: + current_cache["previous_modulated_input"] = self.layers[0].adaLN_modulation(adaln_input.clone())[0] + else: + current_cache = run_data['cache'].get(max_seq_len, {}) + processed_x = x + current_cache.get("previous_residual", 0) + if max_seq_len != run_data.get('uncond_seq_len'): + run_data['results'].setdefault("cache_hits", 0) + run_data['results']["cache_hits"] += 1 + + if num_steps_in_state is not None and max_seq_len != run_data.get('uncond_seq_len'): + run_data['cnt'] += 1 + + output = self.final_layer(processed_x, adaln_input) + if hasattr(self, 'unpatchify'): output = self.unpatchify(output, img_size, cap_size, return_tensor=True)[:, :, :h_img, :w_img] + else: raise NotImplementedError("Model does not have an 'unpatchify' method.") + return -output + + +# TeaCache Patcher +class TeaCache_Patcher: + DEFAULT_COEFFS = [393.76566581, -603.50993606, 209.10239044, -23.00726601, 0.86377344] + + @classmethod + def INPUT_TYPES(cls): + modes = ["手动输入", "自动微调", "贝叶斯优化"] if SKOPT_AVAILABLE else ["手动输入", "自动微调"] + metrics = ["SQ_Ratio (基于命中率)", "SQF_Ratio (基于LPIPS)"] if LPIPS_AVAILABLE else ["SQ_Ratio (基于命中率)"] + return { + "required": { + "model": ("MODEL",), + "mode": (modes,), + "evaluation_metric": (metrics,), + "rel_l1_thresh": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 10.0, "step": 0.001}), + }, + "optional": { + "coefficients_str": ("STRING", {"multiline": True, "default": json.dumps(cls.DEFAULT_COEFFS)}), + } + } + + RETURN_TYPES = ("MODEL", "STRING") + RETURN_NAMES = ("MODEL", "run_id") + FUNCTION = "patch_model" + CATEGORY = "utils/analysis" + + def _load_history(self, file_path): + try: + if os.path.exists(file_path): + with open(file_path, 'r') as f: return json.load(f) + except Exception: pass + return [] + + def _get_score(self, run, baseline_time, metric): + time_saved = baseline_time - run.get('generation_time', 0) + if time_saved <= 0: return 0 + + if "LPIPS" in metric: + # 使用 LPIPS 作为评估 + lpips_dist = run.get("lpips_distance") + if lpips_dist is None or lpips_dist < 1e-6: return 0 + return time_saved / lpips_dist + else: + # 默认使用 Hit Rate 作为评估 + inferences = run.get("total_inferences", 0) + if inferences == 0: return 0 + hit_ratio = run.get("cache_hits", 0) / inferences + if hit_ratio < 1e-6: return 0 + return time_saved * hit_ratio + + def _find_best_run(self, history, baseline_time, metric): + eff_runs = [r for r in history if r.get("rel_l1_thresh") != 0] + if not eff_runs: return None + + best_run, max_score = None, -1.0 + for run in eff_runs: + score = self._get_score(run, baseline_time, metric) + if score > max_score: + max_score, best_run = score, run + return best_run + + def patch_model(self, model, mode, evaluation_metric, rel_l1_thresh, coefficients_str=None): + history = self._load_history(os.path.join(folder_paths.get_output_directory(), "teacache_analysis.json")) + baseline_runs = [r for r in history if r.get("rel_l1_thresh") == 0] + baseline_time = min([r['generation_time'] for r in baseline_runs]) if baseline_runs else None + + coeffs_to_use = [] + if mode == "贝叶斯优化": + if not SKOPT_AVAILABLE: raise Exception("'贝叶斯优化' 模式需要 'scikit-optimize' 库。") + best_efficient_run = self._find_best_run(history, baseline_time, evaluation_metric) + space = [Real(-1000, 1000), Real(-1000, 1000), Real(-1000, 1000), Real(-200, 200), Real(-50, 50)] + optimizer = Optimizer(dimensions=space, random_state=int(time.time()), acq_func="gp_hedge") + + if history and baseline_time: + valid_history = [r for r in history if "coefficients" in r and r.get("rel_l1_thresh") != 0 and len(r["coefficients"]) == len(space)] + if "LPIPS" in evaluation_metric: + valid_history = [r for r in valid_history if r.get("lpips_distance") is not None] + + if valid_history: + y_iters = [-self._get_score(r, baseline_time, evaluation_metric) for r in valid_history] + optimizer.tell([r["coefficients"] for r in valid_history], y_iters) + coeffs_to_use = optimizer.ask() + + elif mode == "自动微调": + best_efficient_run = self._find_best_run(history, baseline_time, evaluation_metric) + base_coeffs = best_efficient_run["coefficients"] if best_efficient_run else self.DEFAULT_COEFFS + idx_to_tweak = len(history) % len(base_coeffs) + perturb_factor = np.random.uniform(0.8, 1.2) + coeffs_to_use = list(base_coeffs) + coeffs_to_use[idx_to_tweak] *= perturb_factor + else: # 手动输入 + coeffs_to_use = json.loads(coefficients_str) if coefficients_str else self.DEFAULT_COEFFS + + run_id = str(time.time_ns()) + params_for_run = { "rel_l1_thresh": rel_l1_thresh, "coefficients_to_use": coeffs_to_use } + STATE_MANAGER.register_run(run_id, params_for_run) + + new_model = model.clone() + diffusion_model = new_model.get_model_object("diffusion_model") + + if not hasattr(diffusion_model, 'forward_original'): + diffusion_model.forward_original = diffusion_model.forward + diffusion_model.forward = teacache_forward_analysis.__get__(diffusion_model, diffusion_model.__class__) + + def unet_wrapper_function(model_function, kwargs): + c_dict = kwargs.get("c", {}) + c_dict.setdefault("transformer_options", {}) + c_dict["transformer_options"]["teacache_run_id"] = run_id + if "sample_sigmas" in c_dict["transformer_options"]: + c_dict["transformer_options"]["num_steps"] = len(c_dict["transformer_options"]["sample_sigmas"]) + return model_function(kwargs["input"], kwargs["timestep"], **c_dict) + + new_model.set_model_unet_function_wrapper(unet_wrapper_function) + print(f"[TeaCache Patcher] 已准备好运行,ID: {run_id}, 评估模式: {evaluation_metric}") + return (new_model, run_id) + +# 结果收集 +class TeaCache_Result_Collector: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "latent": ("LATENT",), + "run_id": ("STRING", {"forceInput": True}), + "analysis_file": ("STRING", {"default": "teacache_analysis.json"}), + }, + "optional": { + "trigger": ("STRING", {"forceInput": True}) + } + } + + RETURN_TYPES = ("LATENT", "STRING") + RETURN_NAMES = ("LATENT", "status") + FUNCTION = "collect_and_save" + CATEGORY = "utils/analysis" + + def _load_history(self, file_path): + try: + if os.path.exists(file_path): + with open(file_path, 'r') as f: return json.load(f) + except Exception: pass + return [] + + def collect_and_save(self, latent, run_id, analysis_file, trigger=None): + run_data = STATE_MANAGER.get_run_data(run_id) + if not run_data: + return (latent, f"错误: 未找到ID为 {run_id} 的运行记录。可能已被过早清理。") + + params = run_data['params'] + results = run_data['results'] + generation_time = time.time() - run_data.get('start_time', time.time()) + + final_run_data = { + "timestamp": time.time(), + "generation_time": generation_time, + "cache_hits": results.get('cache_hits', 0), + "total_inferences": results.get('total_inferences', 0), + "rel_l1_thresh": params.get('rel_l1_thresh'), + "coefficients": params.get('coefficients_to_use'), + "lpips_distance": results.get('lpips_distance', None) + } + + full_path = os.path.join(folder_paths.get_output_directory(), analysis_file) + try: + history_data = self._load_history(full_path) + history_data.append(final_run_data) + with open(full_path, 'w') as f: json.dump(history_data, f, indent=4) + STATE_MANAGER.cleanup_run(run_id) + + summary = f"结果已成功保存到 {analysis_file}。\n" + summary += f"耗时: {generation_time:.2f}s, 缓存命中/总数: {results.get('cache_hits', 0)}/{results.get('total_inferences', 0)}" + if final_run_data["lpips_distance"] is not None: + summary += f", LPIPS距离: {final_run_data['lpips_distance']:.4f}" + print(summary) + return (latent, summary) + except Exception as e: + return (latent, f"写入文件时出错: {e}") + +# LPIPS模型加载 +class LPIPS_Model_Loader: + def __init__(self): + self.model = None + + @classmethod + def INPUT_TYPES(cls): + return {"required": {}} + + RETURN_TYPES = ("LPIPS_MODEL",) + FUNCTION = "load_model" + CATEGORY = "utils/analysis" + + def load_model(self): + if not LPIPS_AVAILABLE: + raise Exception("LPIPS库未安装。请执行 'pip install lpips'") + if STATE_MANAGER.get_lpips_model() is None: + print("正在加载LPIPS模型 (vgg)...") + # 使用 .cpu() 确保模型在CPU上,避免多余的GPU显存占用 + lpips_model = lpips.LPIPS(net='vgg').cpu() + STATE_MANAGER.set_lpips_model(lpips_model) + print("LPIPS模型加载完成。") + return (STATE_MANAGER.get_lpips_model(),) + +# 基准图像存储 +class Store_Baseline_Image: + @classmethod + def INPUT_TYPES(cls): + return {"required": {"image": ("IMAGE",)}} + + RETURN_TYPES = ("BASELINE_IMG",) + FUNCTION = "store_image" + CATEGORY = "utils/analysis" + + def store_image(self, image): + baseline_tensor = image.permute(0, 3, 1, 2).contiguous() + STATE_MANAGER.set_baseline_image(baseline_tensor) + return (baseline_tensor.clone(),) + +# LPIPS评估 +class TeaCache_LPIPS_Evaluator: + @classmethod + def INPUT_TYPES(cls): + return {"required": { + "test_image": ("IMAGE",), + "baseline_image": ("BASELINE_IMG",), + "lpips_model": ("LPIPS_MODEL",), + "run_id": ("STRING", {"forceInput": True}), + }} + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("status",) + FUNCTION = "evaluate" + CATEGORY = "utils/analysis" + + def _preprocess_image(self, image_tensor): + # 将 [0,1] 范围的图像张量转换为 [-1,1] + return image_tensor * 2.0 - 1.0 + + def evaluate(self, test_image, baseline_image, lpips_model, run_id): + if baseline_image is None: + return ("错误: 未提供基准图像 (Baseline Image)。",) + + run_data = STATE_MANAGER.get_run_data(run_id) + if not run_data: + return (f"错误: 未找到ID为 {run_id} 的运行记录。",) + + # 将测试图像从 [B, H, W, C] 转换为 [B, C, H, W] + test_image_t = test_image.permute(0, 3, 1, 2) + + # 预处理 + test_img_proc = self._preprocess_image(test_image_t) + base_img_proc = self._preprocess_image(baseline_image) + + # 计算LPIPS距离 + device = 'cpu' + lpips_model.to(device) + distance = lpips_model(test_img_proc.to(device), base_img_proc.to(device)) + + lpips_score = distance.item() + run_data.setdefault('results', {})['lpips_distance'] = lpips_score + + return (f"LPIPS距离计算完成: {lpips_score:.4f}",)