diff --git a/embedding_adapters.py b/embedding_adapters.py index 03eda9e..73de078 100644 --- a/embedding_adapters.py +++ b/embedding_adapters.py @@ -188,43 +188,43 @@ class GeminiEmbeddingAdapter(BaseEmbeddingAdapter): return [] class SiliconFlowEmbeddingAdapter(BaseEmbeddingAdapter): - """ - 基于 SiliconFlow 的 embedding 适配器 - """ - def __init__(self, api_key: str, base_url: str, model_name: str): - # 自动为 base_url 添加 scheme(如果缺失) - if not base_url.startswith("http://") and not base_url.startswith("https://"): - base_url = "https://" + base_url - self.url = base_url if base_url else "https://api.siliconflow.cn/v1/embeddings" - - self.payload = { - "model": model_name, - "input": "Silicon flow embedding online: fast, affordable, and high-quality embedding services. come try it out!", - "encoding_format": "float" - } - self.headers = { - "Authorization": "Bearer {api_key}".format(api_key=api_key), - "Content-Type": "application/json" - } - - def embed_documents(self, texts: List[str]) -> List[List[float]]: - embeddings = [] - for text in texts: - self.payload["input"] = text - response = requests.post(self.url, json=self.payload, headers=self.headers) - result = response.json() - # 从返回数据中提取第一个 embedding - emb = result.get("data", [{}])[0].get("embedding", []) - embeddings.append(emb) - return embeddings - - def embed_query(self, query: str) -> List[float]: - self.payload["input"] = query - # print('SiliconFlowEmbeddingAdapter发送',self.payload) - response = requests.post(self.url, json=self.payload, headers=self.headers) - result = response.json() - return result.get("data", [{}])[0].get("embedding", []) - + """ + 基于 SiliconFlow 的 embedding 适配器 + """ + def __init__(self, api_key: str, base_url: str, model_name: str): + # 自动为 base_url 添加 scheme(如果缺失) + if not base_url.startswith("http://") and not base_url.startswith("https://"): + base_url = "https://" + base_url + self.url = base_url if base_url else "https://api.siliconflow.cn/v1/embeddings" + + self.payload = { + "model": model_name, + "input": "Silicon flow embedding online: fast, affordable, and high-quality embedding services. come try it out!", + "encoding_format": "float" + } + self.headers = { + "Authorization": "Bearer {api_key}".format(api_key=api_key), + "Content-Type": "application/json" + } + + def embed_documents(self, texts: List[str]) -> List[List[float]]: + embeddings = [] + for text in texts: + self.payload["input"] = text + response = requests.post(self.url, json=self.payload, headers=self.headers) + result = response.json() + # 从返回数据中提取第一个 embedding + emb = result.get("data", [{}])[0].get("embedding", []) + embeddings.append(emb) + return embeddings + + def embed_query(self, query: str) -> List[float]: + self.payload["input"] = query + # print('SiliconFlowEmbeddingAdapter发送',self.payload) + response = requests.post(self.url, json=self.payload, headers=self.headers) + result = response.json() + return result.get("data", [{}])[0].get("embedding", []) + def create_embedding_adapter( interface_format: str, api_key: str, @@ -246,6 +246,6 @@ def create_embedding_adapter( elif fmt == "gemini": return GeminiEmbeddingAdapter(api_key, model_name, base_url) elif fmt == "siliconflow": - return SiliconFlowEmbeddingAdapter(api_key, base_url, model_name) + return SiliconFlowEmbeddingAdapter(api_key, base_url, model_name) else: raise ValueError(f"Unknown embedding interface_format: {interface_format}") diff --git a/ui/config_tab.py b/ui/config_tab.py index c8e5161..2840679 100644 --- a/ui/config_tab.py +++ b/ui/config_tab.py @@ -184,10 +184,9 @@ def build_embeddings_config_tab(self): elif new_value == "Gemini": self.embedding_url_var.set("https://generativelanguage.googleapis.com/v1beta/") self.embedding_model_name_var.set("models/text-embedding-004") - elif new_value == "硅基流动": - self.embedding_url_var.set("https://api.siliconflow.cn/v1/embeddings") - self.embedding_model_name_var.set("BAAI/bge-m3") - + elif new_value == "SiliconFlow": + self.embedding_url_var.set("https://api.siliconflow.cn/v1/embeddings") + self.embedding_model_name_var.set("BAAI/bge-m3") for i in range(5): self.embeddings_config_tab.grid_rowconfigure(i, weight=0) @@ -202,7 +201,9 @@ def build_embeddings_config_tab(self): # 2) Embedding 接口格式 create_label_with_help(self, parent=self.embeddings_config_tab, label_text="Embedding 接口格式:", tooltip_key="embedding_interface_format", row=1, column=0, font=("Microsoft YaHei", 12)) - emb_interface_options = ["DeepSeek", "OpenAI", "Azure OpenAI", "Gemini", "Ollama", "ML Studio","硅基流动"] + + emb_interface_options = ["DeepSeek", "OpenAI", "Azure OpenAI", "Gemini", "Ollama", "ML Studio","SiliconFlow"] + emb_interface_dropdown = ctk.CTkOptionMenu(self.embeddings_config_tab, values=emb_interface_options, variable=self.embedding_interface_format_var, command=on_embedding_interface_changed, font=("Microsoft YaHei", 12)) emb_interface_dropdown.grid(row=1, column=1, padx=5, pady=5, sticky="nsew") diff --git a/ui/novel_params_tab.py b/ui/novel_params_tab.py index 125ed33..84369bb 100644 --- a/ui/novel_params_tab.py +++ b/ui/novel_params_tab.py @@ -4,7 +4,6 @@ import customtkinter as ctk from tkinter import filedialog, messagebox from ui.context_menu import TextWidgetContextMenu from tooltips import tooltips -from tooltips import tooltips def build_novel_params_area(self, start_row=1): self.params_frame = ctk.CTkScrollableFrame(self.right_frame, orientation="vertical")