diff --git a/consistency_checker.py b/consistency_checker.py index 7e6abe3..a205764 100644 --- a/consistency_checker.py +++ b/consistency_checker.py @@ -1,6 +1,6 @@ # consistency_checker.py # -*- coding: utf-8 -*- -from langchain_openai import ChatOpenAI +from llm_adapters import create_llm_adapter # ============== 增加对“剧情要点/未解决冲突”进行检查的可选引导 ============== CONSISTENCY_PROMPT = """\ @@ -32,7 +32,10 @@ def check_consistency( base_url: str, model_name: str, temperature: float = 0.3, - plot_arcs: str = "" # 新增参数,默认空字符串 + plot_arcs: str = "", + interface_format: str = "OpenAI", + max_tokens: int = 2048, + timeout: int = 600 ) -> str: """ 调用模型做简单的一致性检查。可扩展更多提示或校验规则。 @@ -45,20 +48,25 @@ def check_consistency( plot_arcs=plot_arcs, chapter_text=chapter_text ) - model = ChatOpenAI( - model=model_name, - api_key=api_key, + + llm_adapter = create_llm_adapter( + interface_format=interface_format, base_url=base_url, - temperature=temperature + model_name=model_name, + api_key=api_key, + temperature=temperature, + max_tokens=max_tokens, + timeout=timeout ) + # 调试日志 print("\n[ConsistencyChecker] Prompt >>>", prompt) - response = model.invoke(prompt) + response = llm_adapter.invoke(prompt) if not response: return "审校Agent无回复" - + # 调试日志 - print("[ConsistencyChecker] Response <<<", response.content.strip()) + print("[ConsistencyChecker] Response <<<", response) - return response.content.strip() + return response diff --git a/ui.py b/ui.py index ec9a2a2..5d5b36c 100644 --- a/ui.py +++ b/ui.py @@ -1083,6 +1083,9 @@ class NovelGeneratorGUI: base_url = self.base_url_var.get().strip() model_name = self.model_name_var.get().strip() temperature = self.temperature_var.get() + interface_format = self.interface_format_var.get() + max_tokens = self.max_tokens_var.get() + timeout = self.timeout_var.get() chap_num = self.safe_get_int(self.chapter_num_var, 1) chap_file = os.path.join(filepath, "chapters", f"chapter_{chap_num}.txt") @@ -1102,6 +1105,9 @@ class NovelGeneratorGUI: base_url=base_url, model_name=model_name, temperature=temperature, + interface_format=interface_format, + max_tokens=max_tokens, + timeout=timeout, plot_arcs="" ) self.safe_log("审校结果:")