修复贡献于[TdRoseval](https://github.com/TdRoseval)的Pr: #97 所遗留问题;
This commit is contained in:
YILING0013
2025-02-12 10:49:28 +08:00
parent 8ed0166d3e
commit 186566f173
2 changed files with 45 additions and 20 deletions
+40 -19
View File
@@ -135,34 +135,56 @@ class MLStudioEmbeddingAdapter(BaseEmbeddingAdapter):
class GeminiEmbeddingAdapter(BaseEmbeddingAdapter): class GeminiEmbeddingAdapter(BaseEmbeddingAdapter):
""" """
基于 Google Generative AI Gemini接口的 Embedding 适配器 基于 Google Generative AI (Gemini) 接口的 Embedding 适配器
使用直接 POST 请求方式,URL 示例:
https://generativelanguage.googleapis.com/v1beta/models/text-embedding-004:embedContent?key=YOUR_API_KEY
""" """
def __init__(self, api_key: str, model_name: str): def __init__(self, api_key: str, model_name: str, base_url: str):
from google import genai """
# 全局配置,也可根据需要改成 Client(...) 初始化方式 :param api_key: 传入的 Google API Key
genai.configure(api_key=api_key) :param model_name: 这里一般是 "text-embedding-004"
:param base_url: e.g. https://generativelanguage.googleapis.com/v1beta/models
"""
self.api_key = api_key
self.model_name = model_name self.model_name = model_name
self.base_url = base_url.rstrip("/")
def embed_documents(self, texts: List[str]) -> List[List[float]]: def embed_documents(self, texts: List[str]) -> List[List[float]]:
from google import genai
embeddings = [] embeddings = []
for text in texts: for text in texts:
try: vec = self._embed_single(text)
result = genai.embed_content(model=self.model_name, content=text) embeddings.append(vec)
# 返回结构中包含 'embedding' 字段
embeddings.append(result.get('embedding', []))
except Exception as e:
logging.error(f"Gemini embed_content error: {e}")
embeddings.append([])
return embeddings return embeddings
def embed_query(self, query: str) -> List[float]: def embed_query(self, query: str) -> List[float]:
from google import genai return self._embed_single(query)
def _embed_single(self, text: str) -> List[float]:
"""
直接调用 Google Generative Language API (Gemini) 接口,获取文本 embedding
"""
url = f"{self.base_url}/{self.model_name}:embedContent?key={self.api_key}"
payload = {
"model": self.model_name,
"content": {
"parts": [
{"text": text}
]
}
}
try: try:
result = genai.embed_content(model=self.model_name, content=query) response = requests.post(url, json=payload)
return result.get('embedding', []) print(response.text)
response.raise_for_status()
result = response.json()
embedding_data = result.get("embedding", {})
return embedding_data.get("values", [])
except requests.exceptions.RequestException as e:
logging.error(f"Gemini embed_content request error: {e}\n{traceback.format_exc()}")
return []
except Exception as e: except Exception as e:
logging.error(f"Gemini embed_content error: {e}") logging.error(f"Gemini embed_content parse error: {e}\n{traceback.format_exc()}")
return [] return []
def create_embedding_adapter( def create_embedding_adapter(
@@ -184,7 +206,6 @@ def create_embedding_adapter(
elif fmt == "ml studio": elif fmt == "ml studio":
return MLStudioEmbeddingAdapter(api_key, base_url, model_name) return MLStudioEmbeddingAdapter(api_key, base_url, model_name)
elif fmt == "gemini": elif fmt == "gemini":
# base_url 对 Gemini 暂无用处,可忽略 return GeminiEmbeddingAdapter(api_key, model_name, base_url)
return GeminiEmbeddingAdapter(api_key, model_name)
else: else:
raise ValueError(f"Unknown embedding interface_format: {interface_format}") raise ValueError(f"Unknown embedding interface_format: {interface_format}")
+5 -1
View File
@@ -548,10 +548,14 @@ 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")
self.embedding_model_name_var.set("text-embedding-ada-002")
elif new_value == "Azure OpenAI": elif new_value == "Azure OpenAI":
self.embedding_url_var.set("https://[az].openai.azure.com/openai/deployments/[model]/embeddings?api-version=2023-05-15") 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")
elif new_value == "Gemini":
self.embedding_url_var.set("https://generativelanguage.googleapis.com/v1beta/")
self.embedding_model_name_var.set("models/text-embedding-004")
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)
@@ -580,7 +584,7 @@ class NovelGeneratorGUI:
column=0, column=0,
font=("Microsoft YaHei", 12) font=("Microsoft YaHei", 12)
) )
emb_interface_options = ["DeepSeek", "OpenAI", "Azure OpenAI", "Ollama", "ML Studio"] emb_interface_options = ["DeepSeek", "OpenAI", "Azure OpenAI", "Gemini", "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,