From 2433b6c159bdfe0d394af756bf1966a0430c87a6 Mon Sep 17 00:00:00 2001 From: spawner Date: Sat, 7 Jun 2025 01:07:52 +0800 Subject: [PATCH] Update tc_test.py --- tc_test.py | 116 +++++++++++++++-------------------------------------- 1 file changed, 32 insertions(+), 84 deletions(-) diff --git a/tc_test.py b/tc_test.py index ee4c5f4..2f284ca 100644 --- a/tc_test.py +++ b/tc_test.py @@ -54,7 +54,7 @@ class TeaCacheStateManager: 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') @@ -136,7 +136,6 @@ def teacache_forward_analysis(self, x, timesteps, context, num_tokens, **kwargs) 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] @@ -144,9 +143,8 @@ class TeaCache_Patcher: @classmethod def INPUT_TYPES(cls): modes = ["手动输入", "自动微调", "贝叶斯优化"] if SKOPT_AVAILABLE else ["手动输入", "自动微调"] - metrics = ["SQ_Ratio (基于命中率)", "SQF_Ratio (基于LPIPS)"] if LPIPS_AVAILABLE else ["SQ_Ratio (基于命中率)"] + metrics = ["速度-命中率权衡", "质量-命中率权衡 (LPIPS)"] if LPIPS_AVAILABLE else ["速度-命中率权衡"] - # 只有在LPIPS可用时,才显示质量阈值选项 optional_inputs = { "coefficients_str": ("STRING", {"multiline": True, "default": json.dumps(cls.DEFAULT_COEFFS)}), } @@ -176,43 +174,41 @@ class TeaCache_Patcher: return [] def _get_score(self, run, baseline_time, metric, max_lpips_thresh): - time_saved = baseline_time - run.get('generation_time', 0) - if time_saved <= 0: return 0 + inferences = run.get("total_inferences", 0) + if inferences == 0: return 0 + hit_ratio = run.get("cache_hits", 0) / inferences if "LPIPS" in metric: lpips_dist = run.get("lpips_distance") if lpips_dist is None: return 0 - # 如果LPIPS距离超过了可接受的上限,直接给0分 - # 如果 max_lpips_thresh <= 0,则禁用此功能 if max_lpips_thresh > 0 and lpips_dist > max_lpips_thresh: return 0 - if lpips_dist < 1e-6: return 0 # 避免除以0 - return time_saved / lpips_dist - else: - inferences = run.get("total_inferences", 0) - if inferences == 0: return 0 - hit_ratio = run.get("cache_hits", 0) / inferences + if lpips_dist < 1e-6: return 0 + + # 公式:命中率 / LPIPS距离 + return hit_ratio / lpips_dist + else: # "速度-命中率权衡" + time_saved = baseline_time - run.get('generation_time', 0) + if time_saved <= 0: return 0 return time_saved * hit_ratio - # ✨ 增加 max_lpips_thresh 参数 def _find_best_run(self, history, baseline_time, metric, max_lpips_thresh): 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, max_lpips_thresh) if score > max_score: max_score, best_run = score, run return best_run - # ✨ 增加 max_lpips_thresh 参数 def patch_model(self, model, mode, evaluation_metric, rel_l1_thresh, coefficients_str=None, max_lpips_thresh=0.6): 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 = [] @@ -222,19 +218,20 @@ class TeaCache_Patcher: 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: + if history: 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] - + elif "速度" in evaluation_metric and not baseline_time: + print("警告: 缺少基准时间,无法进行基于速度的贝叶斯优化。") + valid_history = [] + if valid_history: - # ✨ 将阈值传递给打分函数 y_iters = [-self._get_score(r, baseline_time, evaluation_metric, max_lpips_thresh) 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, max_lpips_thresh) base_coeffs = best_efficient_run["coefficients"] if best_efficient_run else self.DEFAULT_COEFFS idx_to_tweak = len(history) % len(base_coeffs) @@ -245,7 +242,6 @@ class TeaCache_Patcher: 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, @@ -276,56 +272,35 @@ class TeaCache_Patcher: 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 { "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} 的运行记录。可能已被过早清理。") - + 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), - "max_lpips_thresh": params.get('max_lpips_thresh', None) # ✨ 保存阈值到JSON + "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), "max_lpips_thresh": params.get('max_lpips_thresh', 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: @@ -337,20 +312,14 @@ class TeaCache_Result_Collector: # LPIPS模型加载 class LPIPS_Model_Loader: - def __init__(self): - self.model = None - + def __init__(self): self.model = None @classmethod - def INPUT_TYPES(cls): - return {"required": {}} - + 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 not LPIPS_AVAILABLE: raise Exception("LPIPS库未安装。请执行 'pip install lpips'") if STATE_MANAGER.get_lpips_model() is None: print("正在加载LPIPS模型 (vgg)...") lpips_model = lpips.LPIPS(net='vgg').cpu() @@ -361,13 +330,10 @@ class LPIPS_Model_Loader: # 基准图像存储 class Store_Baseline_Image: @classmethod - def INPUT_TYPES(cls): - return {"required": {"image": ("IMAGE",)}} - + 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) @@ -376,40 +342,22 @@ class Store_Baseline_Image: # 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}), - }} - + 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): - return image_tensor * 2.0 - 1.0 - + def _preprocess_image(self, image_tensor): 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)。",) - + 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} 的运行记录。",) - + if not run_data: return (f"错误: 未找到ID为 {run_id} 的运行记录。",) 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) - 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}",)