Update tc_test.py

This commit is contained in:
spawner
2025-06-07 01:07:52 +08:00
committed by GitHub
parent 167b5c4f9b
commit 2433b6c159
+32 -84
View File
@@ -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}",)