增加支持SiliconFlow的embedding接口

This commit is contained in:
FynnCX
2025-03-06 00:07:17 +08:00
parent b279f1de80
commit 4f24943923
2 changed files with 52 additions and 4 deletions
+44 -2
View File
@@ -1,10 +1,12 @@
# embedding_adapters.py
# -*- coding: utf-8 -*-
import logging
import requests
import traceback
from typing import List
from langchain_openai import OpenAIEmbeddings, AzureOpenAIEmbeddings
import requests
from langchain_openai import AzureOpenAIEmbeddings, OpenAIEmbeddings
def ensure_openai_base_url_has_v1(url: str) -> str:
"""
@@ -187,6 +189,44 @@ class GeminiEmbeddingAdapter(BaseEmbeddingAdapter):
logging.error(f"Gemini embed_content parse error: {e}\n{traceback.format_exc()}")
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", [])
def create_embedding_adapter(
interface_format: str,
api_key: str,
@@ -207,5 +247,7 @@ def create_embedding_adapter(
return MLStudioEmbeddingAdapter(api_key, base_url, model_name)
elif fmt == "gemini":
return GeminiEmbeddingAdapter(api_key, model_name, base_url)
elif fmt == "siliconflow":
return SiliconFlowEmbeddingAdapter(api_key, base_url, model_name)
else:
raise ValueError(f"Unknown embedding interface_format: {interface_format}")
+8 -2
View File
@@ -1,10 +1,13 @@
# ui/config_tab.py
# -*- coding: utf-8 -*-
import customtkinter as ctk
from tkinter import messagebox
import customtkinter as ctk
from config_manager import load_config, save_config
from tooltips import tooltips
def create_label_with_help(self, parent, label_text, tooltip_key, row, column,
font=None, sticky="e", padx=5, pady=5):
"""
@@ -177,6 +180,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 == "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)
@@ -191,7 +197,7 @@ 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")