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/embedding_adapters.py b/embedding_adapters.py index 9b0d7a8..4304644 100644 --- a/embedding_adapters.py +++ b/embedding_adapters.py @@ -4,7 +4,7 @@ import logging import requests import traceback from typing import List -from langchain_openai import OpenAIEmbeddings +from langchain_openai import OpenAIEmbeddings, AzureOpenAIEmbeddings def ensure_openai_base_url_has_v1(url: str) -> str: """ @@ -45,6 +45,33 @@ class OpenAIEmbeddingAdapter(BaseEmbeddingAdapter): def embed_query(self, query: str) -> List[float]: return self._embedding.embed_query(query) + +class AzureOpenAIEmbeddingAdapter(BaseEmbeddingAdapter): + """ + 基于 AzureOpenAIEmbeddings(或兼容接口)的适配器 + """ + def __init__(self, api_key: str, base_url: str, model_name: str): + import re + match = re.match(r'https://(.+?)/openai/deployments/(.+?)/embeddings\?api-version=(.+)', base_url) + if match: + self.azure_endpoint = f"https://{match.group(1)}" + self.azure_deployment = match.group(2) + self.api_version = match.group(3) + else: + raise ValueError("Invalid Azure OpenAI base_url format") + + self._embedding = AzureOpenAIEmbeddings( + azure_endpoint=self.azure_endpoint, + azure_deployment=self.azure_deployment, + openai_api_key=api_key, + api_version=self.api_version, + ) + + def embed_documents(self, texts: List[str]) -> List[List[float]]: + return self._embedding.embed_documents(texts) + + def embed_query(self, query: str) -> List[float]: + return self._embedding.embed_query(query) class OllamaEmbeddingAdapter(BaseEmbeddingAdapter): """ @@ -112,6 +139,8 @@ def create_embedding_adapter( """ if interface_format.lower() == "openai": return OpenAIEmbeddingAdapter(api_key, base_url, model_name) + elif interface_format.lower() == "azure openai": + return AzureOpenAIEmbeddingAdapter(api_key, base_url, model_name) elif interface_format.lower() == "ollama": return OllamaEmbeddingAdapter(model_name, base_url) elif interface_format.lower() == "ml studio": diff --git a/llm_adapters.py b/llm_adapters.py index 5bc8fc9..f09d06b 100644 --- a/llm_adapters.py +++ b/llm_adapters.py @@ -2,7 +2,7 @@ # -*- coding: utf-8 -*- import logging from typing import Optional -from langchain_openai import ChatOpenAI +from langchain_openai import ChatOpenAI, AzureChatOpenAI def ensure_openai_base_url_has_v1(url: str) -> str: import re @@ -77,6 +77,43 @@ class OpenAIAdapter(BaseLLMAdapter): return "" return response.content +class AzureOpenAIAdapter(BaseLLMAdapter): + """ + 适配 Azure OpenAI 接口(使用 langchain.ChatOpenAI) + """ + def __init__(self, api_key: str, base_url: str, model_name: str, max_tokens: int, temperature: float = 0.7, timeout: Optional[int] = 600): + import re + match = re.match(r'https://(.+?)/openai/deployments/(.+?)/chat/completions\?api-version=(.+)', base_url) + if match: + self.azure_endpoint = f"https://{match.group(1)}" + self.azure_deployment = match.group(2) + self.api_version = match.group(3) + else: + raise ValueError("Invalid Azure OpenAI base_url format") + + self.api_key = api_key + self.model_name = self.azure_deployment + self.max_tokens = max_tokens + self.temperature = temperature + self.timeout = timeout + + self._client = AzureChatOpenAI( + azure_endpoint=self.azure_endpoint, + azure_deployment=self.azure_deployment, + api_version=self.api_version, + api_key=self.api_key, + max_tokens=self.max_tokens, + temperature=self.temperature, + timeout=self.timeout + ) + + def invoke(self, prompt: str) -> str: + response = self._client.invoke(prompt) + if not response: + logging.warning("No response from AzureOpenAIAdapter.") + return "" + return response.content + class OllamaAdapter(BaseLLMAdapter): """ Ollama 同样有一个 OpenAI-like /v1/chat 接口,可直接使用 ChatOpenAI。 @@ -147,6 +184,8 @@ def create_llm_adapter( return DeepSeekAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) elif interface_format.lower() == "openai": return OpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) + elif interface_format.lower() == "azure openai": + return AzureOpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) elif interface_format.lower() == "ollama": return OllamaAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) elif interface_format.lower() == "ml studio": diff --git a/ui.py b/ui.py index a9845b8..5d5b36c 100644 --- a/ui.py +++ b/ui.py @@ -290,6 +290,8 @@ class NovelGeneratorGUI: self.base_url_var.set("http://localhost:1234/v1") elif new_value == "OpenAI": self.base_url_var.set("https://api.openai.com/v1") + elif new_value == "Azure OpenAI": + self.base_url_var.set("https://[az].openai.azure.com/openai/deployments/[model]/chat/completions?api-version=2024-08-01-preview") elif new_value == "DeepSeek": self.base_url_var.set("https://api.deepseek.com/v1") @@ -332,7 +334,7 @@ class NovelGeneratorGUI: column=0, font=("Microsoft YaHei", 12) ) - interface_options = ["DeepSeek", "OpenAI", "Ollama", "ML Studio"] + interface_options = ["DeepSeek", "OpenAI", "Azure OpenAI", "Ollama", "ML Studio"] interface_dropdown = ctk.CTkOptionMenu( self.ai_config_tab, values=interface_options, @@ -452,6 +454,8 @@ class NovelGeneratorGUI: self.embedding_url_var.set("http://localhost:1234/v1") elif new_value == "OpenAI": self.embedding_url_var.set("https://api.openai.com/v1") + elif new_value == "Azure OpenAI": + self.embedding_url_var.set("https://[az].openai.azure.com/openai/deployments/[model]/embeddings?api-version=2023-05-15") elif new_value == "DeepSeek": self.embedding_url_var.set("https://api.deepseek.com/v1") @@ -482,7 +486,7 @@ class NovelGeneratorGUI: column=0, font=("Microsoft YaHei", 12) ) - emb_interface_options = ["DeepSeek", "OpenAI", "Ollama", "ML Studio"] + emb_interface_options = ["DeepSeek", "OpenAI", "Azure OpenAI", "Ollama", "ML Studio"] emb_interface_dropdown = ctk.CTkOptionMenu( self.embeddings_config_tab, values=emb_interface_options, @@ -1079,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") @@ -1098,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("审校结果:")