Merge branch 'main' into main
This commit is contained in:
+38
-38
@@ -188,43 +188,43 @@ class GeminiEmbeddingAdapter(BaseEmbeddingAdapter):
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
class SiliconFlowEmbeddingAdapter(BaseEmbeddingAdapter):
|
class SiliconFlowEmbeddingAdapter(BaseEmbeddingAdapter):
|
||||||
"""
|
"""
|
||||||
基于 SiliconFlow 的 embedding 适配器
|
基于 SiliconFlow 的 embedding 适配器
|
||||||
"""
|
"""
|
||||||
def __init__(self, api_key: str, base_url: str, model_name: str):
|
def __init__(self, api_key: str, base_url: str, model_name: str):
|
||||||
# 自动为 base_url 添加 scheme(如果缺失)
|
# 自动为 base_url 添加 scheme(如果缺失)
|
||||||
if not base_url.startswith("http://") and not base_url.startswith("https://"):
|
if not base_url.startswith("http://") and not base_url.startswith("https://"):
|
||||||
base_url = "https://" + base_url
|
base_url = "https://" + base_url
|
||||||
self.url = base_url if base_url else "https://api.siliconflow.cn/v1/embeddings"
|
self.url = base_url if base_url else "https://api.siliconflow.cn/v1/embeddings"
|
||||||
|
|
||||||
self.payload = {
|
self.payload = {
|
||||||
"model": model_name,
|
"model": model_name,
|
||||||
"input": "Silicon flow embedding online: fast, affordable, and high-quality embedding services. come try it out!",
|
"input": "Silicon flow embedding online: fast, affordable, and high-quality embedding services. come try it out!",
|
||||||
"encoding_format": "float"
|
"encoding_format": "float"
|
||||||
}
|
}
|
||||||
self.headers = {
|
self.headers = {
|
||||||
"Authorization": "Bearer {api_key}".format(api_key=api_key),
|
"Authorization": "Bearer {api_key}".format(api_key=api_key),
|
||||||
"Content-Type": "application/json"
|
"Content-Type": "application/json"
|
||||||
}
|
}
|
||||||
|
|
||||||
def embed_documents(self, texts: List[str]) -> List[List[float]]:
|
def embed_documents(self, texts: List[str]) -> List[List[float]]:
|
||||||
embeddings = []
|
embeddings = []
|
||||||
for text in texts:
|
for text in texts:
|
||||||
self.payload["input"] = text
|
self.payload["input"] = text
|
||||||
response = requests.post(self.url, json=self.payload, headers=self.headers)
|
response = requests.post(self.url, json=self.payload, headers=self.headers)
|
||||||
result = response.json()
|
result = response.json()
|
||||||
# 从返回数据中提取第一个 embedding
|
# 从返回数据中提取第一个 embedding
|
||||||
emb = result.get("data", [{}])[0].get("embedding", [])
|
emb = result.get("data", [{}])[0].get("embedding", [])
|
||||||
embeddings.append(emb)
|
embeddings.append(emb)
|
||||||
return embeddings
|
return embeddings
|
||||||
|
|
||||||
def embed_query(self, query: str) -> List[float]:
|
def embed_query(self, query: str) -> List[float]:
|
||||||
self.payload["input"] = query
|
self.payload["input"] = query
|
||||||
# print('SiliconFlowEmbeddingAdapter发送',self.payload)
|
# print('SiliconFlowEmbeddingAdapter发送',self.payload)
|
||||||
response = requests.post(self.url, json=self.payload, headers=self.headers)
|
response = requests.post(self.url, json=self.payload, headers=self.headers)
|
||||||
result = response.json()
|
result = response.json()
|
||||||
return result.get("data", [{}])[0].get("embedding", [])
|
return result.get("data", [{}])[0].get("embedding", [])
|
||||||
|
|
||||||
def create_embedding_adapter(
|
def create_embedding_adapter(
|
||||||
interface_format: str,
|
interface_format: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
@@ -246,6 +246,6 @@ def create_embedding_adapter(
|
|||||||
elif fmt == "gemini":
|
elif fmt == "gemini":
|
||||||
return GeminiEmbeddingAdapter(api_key, model_name, base_url)
|
return GeminiEmbeddingAdapter(api_key, model_name, base_url)
|
||||||
elif fmt == "siliconflow":
|
elif fmt == "siliconflow":
|
||||||
return SiliconFlowEmbeddingAdapter(api_key, base_url, model_name)
|
return SiliconFlowEmbeddingAdapter(api_key, base_url, model_name)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unknown embedding interface_format: {interface_format}")
|
raise ValueError(f"Unknown embedding interface_format: {interface_format}")
|
||||||
|
|||||||
+6
-5
@@ -184,10 +184,9 @@ def build_embeddings_config_tab(self):
|
|||||||
elif new_value == "Gemini":
|
elif new_value == "Gemini":
|
||||||
self.embedding_url_var.set("https://generativelanguage.googleapis.com/v1beta/")
|
self.embedding_url_var.set("https://generativelanguage.googleapis.com/v1beta/")
|
||||||
self.embedding_model_name_var.set("models/text-embedding-004")
|
self.embedding_model_name_var.set("models/text-embedding-004")
|
||||||
elif new_value == "硅基流动":
|
elif new_value == "SiliconFlow":
|
||||||
self.embedding_url_var.set("https://api.siliconflow.cn/v1/embeddings")
|
self.embedding_url_var.set("https://api.siliconflow.cn/v1/embeddings")
|
||||||
self.embedding_model_name_var.set("BAAI/bge-m3")
|
self.embedding_model_name_var.set("BAAI/bge-m3")
|
||||||
|
|
||||||
|
|
||||||
for i in range(5):
|
for i in range(5):
|
||||||
self.embeddings_config_tab.grid_rowconfigure(i, weight=0)
|
self.embeddings_config_tab.grid_rowconfigure(i, weight=0)
|
||||||
@@ -202,7 +201,9 @@ def build_embeddings_config_tab(self):
|
|||||||
|
|
||||||
# 2) Embedding 接口格式
|
# 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))
|
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 = 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")
|
emb_interface_dropdown.grid(row=1, column=1, padx=5, pady=5, sticky="nsew")
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import customtkinter as ctk
|
|||||||
from tkinter import filedialog, messagebox
|
from tkinter import filedialog, messagebox
|
||||||
from ui.context_menu import TextWidgetContextMenu
|
from ui.context_menu import TextWidgetContextMenu
|
||||||
from tooltips import tooltips
|
from tooltips import tooltips
|
||||||
from tooltips import tooltips
|
|
||||||
|
|
||||||
def build_novel_params_area(self, start_row=1):
|
def build_novel_params_area(self, start_row=1):
|
||||||
self.params_frame = ctk.CTkScrollableFrame(self.right_frame, orientation="vertical")
|
self.params_frame = ctk.CTkScrollableFrame(self.right_frame, orientation="vertical")
|
||||||
|
|||||||
Reference in New Issue
Block a user