一致性检查器使用 llm_adapters

This commit is contained in:
桑榆肖物
2025-02-09 23:34:07 +08:00
parent 8afa1083e0
commit 20805abbd1
2 changed files with 24 additions and 10 deletions
+18 -10
View File
@@ -1,6 +1,6 @@
# consistency_checker.py # consistency_checker.py
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from langchain_openai import ChatOpenAI from llm_adapters import create_llm_adapter
# ============== 增加对“剧情要点/未解决冲突”进行检查的可选引导 ============== # ============== 增加对“剧情要点/未解决冲突”进行检查的可选引导 ==============
CONSISTENCY_PROMPT = """\ CONSISTENCY_PROMPT = """\
@@ -32,7 +32,10 @@ def check_consistency(
base_url: str, base_url: str,
model_name: str, model_name: str,
temperature: float = 0.3, temperature: float = 0.3,
plot_arcs: str = "" # 新增参数,默认空字符串 plot_arcs: str = "",
interface_format: str = "OpenAI",
max_tokens: int = 2048,
timeout: int = 600
) -> str: ) -> str:
""" """
调用模型做简单的一致性检查。可扩展更多提示或校验规则。 调用模型做简单的一致性检查。可扩展更多提示或校验规则。
@@ -45,20 +48,25 @@ def check_consistency(
plot_arcs=plot_arcs, plot_arcs=plot_arcs,
chapter_text=chapter_text chapter_text=chapter_text
) )
model = ChatOpenAI(
model=model_name, llm_adapter = create_llm_adapter(
api_key=api_key, interface_format=interface_format,
base_url=base_url, 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) print("\n[ConsistencyChecker] Prompt >>>", prompt)
response = model.invoke(prompt) response = llm_adapter.invoke(prompt)
if not response: if not response:
return "审校Agent无回复" return "审校Agent无回复"
# 调试日志 # 调试日志
print("[ConsistencyChecker] Response <<<", response.content.strip()) print("[ConsistencyChecker] Response <<<", response)
return response.content.strip() return response
+6
View File
@@ -1083,6 +1083,9 @@ class NovelGeneratorGUI:
base_url = self.base_url_var.get().strip() base_url = self.base_url_var.get().strip()
model_name = self.model_name_var.get().strip() model_name = self.model_name_var.get().strip()
temperature = self.temperature_var.get() 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_num = self.safe_get_int(self.chapter_num_var, 1)
chap_file = os.path.join(filepath, "chapters", f"chapter_{chap_num}.txt") chap_file = os.path.join(filepath, "chapters", f"chapter_{chap_num}.txt")
@@ -1102,6 +1105,9 @@ class NovelGeneratorGUI:
base_url=base_url, base_url=base_url,
model_name=model_name, model_name=model_name,
temperature=temperature, temperature=temperature,
interface_format=interface_format,
max_tokens=max_tokens,
timeout=timeout,
plot_arcs="" plot_arcs=""
) )
self.safe_log("审校结果:") self.safe_log("审校结果:")