Update tc_test.py
This commit is contained in:
+32
-84
@@ -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}",)
|
||||
|
||||
Reference in New Issue
Block a user