Merge pull request #89 from sangyuxiaowu/azure
添加对 Azure OpenAI 的支持,新增适配器并更新界面选项
This commit is contained in:
+17
-9
@@ -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
|
||||||
|
|||||||
+30
-1
@@ -4,7 +4,7 @@ import logging
|
|||||||
import requests
|
import requests
|
||||||
import traceback
|
import traceback
|
||||||
from typing import List
|
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:
|
def ensure_openai_base_url_has_v1(url: str) -> str:
|
||||||
"""
|
"""
|
||||||
@@ -46,6 +46,33 @@ class OpenAIEmbeddingAdapter(BaseEmbeddingAdapter):
|
|||||||
def embed_query(self, query: str) -> List[float]:
|
def embed_query(self, query: str) -> List[float]:
|
||||||
return self._embedding.embed_query(query)
|
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):
|
class OllamaEmbeddingAdapter(BaseEmbeddingAdapter):
|
||||||
"""
|
"""
|
||||||
其接口路径为 /api/embeddings
|
其接口路径为 /api/embeddings
|
||||||
@@ -112,6 +139,8 @@ def create_embedding_adapter(
|
|||||||
"""
|
"""
|
||||||
if interface_format.lower() == "openai":
|
if interface_format.lower() == "openai":
|
||||||
return OpenAIEmbeddingAdapter(api_key, base_url, model_name)
|
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":
|
elif interface_format.lower() == "ollama":
|
||||||
return OllamaEmbeddingAdapter(model_name, base_url)
|
return OllamaEmbeddingAdapter(model_name, base_url)
|
||||||
elif interface_format.lower() == "ml studio":
|
elif interface_format.lower() == "ml studio":
|
||||||
|
|||||||
+40
-1
@@ -2,7 +2,7 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
import logging
|
import logging
|
||||||
from typing import Optional
|
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:
|
def ensure_openai_base_url_has_v1(url: str) -> str:
|
||||||
import re
|
import re
|
||||||
@@ -77,6 +77,43 @@ class OpenAIAdapter(BaseLLMAdapter):
|
|||||||
return ""
|
return ""
|
||||||
return response.content
|
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):
|
class OllamaAdapter(BaseLLMAdapter):
|
||||||
"""
|
"""
|
||||||
Ollama 同样有一个 OpenAI-like /v1/chat 接口,可直接使用 ChatOpenAI。
|
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)
|
return DeepSeekAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
|
||||||
elif interface_format.lower() == "openai":
|
elif interface_format.lower() == "openai":
|
||||||
return OpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
|
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":
|
elif interface_format.lower() == "ollama":
|
||||||
return OllamaAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
|
return OllamaAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
|
||||||
elif interface_format.lower() == "ml studio":
|
elif interface_format.lower() == "ml studio":
|
||||||
|
|||||||
@@ -290,6 +290,8 @@ class NovelGeneratorGUI:
|
|||||||
self.base_url_var.set("http://localhost:1234/v1")
|
self.base_url_var.set("http://localhost:1234/v1")
|
||||||
elif new_value == "OpenAI":
|
elif new_value == "OpenAI":
|
||||||
self.base_url_var.set("https://api.openai.com/v1")
|
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":
|
elif new_value == "DeepSeek":
|
||||||
self.base_url_var.set("https://api.deepseek.com/v1")
|
self.base_url_var.set("https://api.deepseek.com/v1")
|
||||||
|
|
||||||
@@ -332,7 +334,7 @@ class NovelGeneratorGUI:
|
|||||||
column=0,
|
column=0,
|
||||||
font=("Microsoft YaHei", 12)
|
font=("Microsoft YaHei", 12)
|
||||||
)
|
)
|
||||||
interface_options = ["DeepSeek", "OpenAI", "Ollama", "ML Studio"]
|
interface_options = ["DeepSeek", "OpenAI", "Azure OpenAI", "Ollama", "ML Studio"]
|
||||||
interface_dropdown = ctk.CTkOptionMenu(
|
interface_dropdown = ctk.CTkOptionMenu(
|
||||||
self.ai_config_tab,
|
self.ai_config_tab,
|
||||||
values=interface_options,
|
values=interface_options,
|
||||||
@@ -452,6 +454,8 @@ class NovelGeneratorGUI:
|
|||||||
self.embedding_url_var.set("http://localhost:1234/v1")
|
self.embedding_url_var.set("http://localhost:1234/v1")
|
||||||
elif new_value == "OpenAI":
|
elif new_value == "OpenAI":
|
||||||
self.embedding_url_var.set("https://api.openai.com/v1")
|
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":
|
elif new_value == "DeepSeek":
|
||||||
self.embedding_url_var.set("https://api.deepseek.com/v1")
|
self.embedding_url_var.set("https://api.deepseek.com/v1")
|
||||||
|
|
||||||
@@ -482,7 +486,7 @@ class NovelGeneratorGUI:
|
|||||||
column=0,
|
column=0,
|
||||||
font=("Microsoft YaHei", 12)
|
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(
|
emb_interface_dropdown = ctk.CTkOptionMenu(
|
||||||
self.embeddings_config_tab,
|
self.embeddings_config_tab,
|
||||||
values=emb_interface_options,
|
values=emb_interface_options,
|
||||||
@@ -1079,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")
|
||||||
@@ -1098,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("审校结果:")
|
||||||
|
|||||||
Reference in New Issue
Block a user