修复贡献于[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):
"""
基于 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):
from google import genai
# 全局配置,也可根据需要改成 Client(...) 初始化方式
genai.configure(api_key=api_key)
def __init__(self, api_key: str, model_name: str, base_url: str):
"""
:param api_key: 传入的 Google 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.base_url = base_url.rstrip("/")
def embed_documents(self, texts: List[str]) -> List[List[float]]:
from google import genai
embeddings = []
for text in texts:
try:
result = genai.embed_content(model=self.model_name, content=text)
# 返回结构中包含 'embedding' 字段
embeddings.append(result.get('embedding', []))
except Exception as e:
logging.error(f"Gemini embed_content error: {e}")
embeddings.append([])
vec = self._embed_single(text)
embeddings.append(vec)
return embeddings
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:
result = genai.embed_content(model=self.model_name, content=query)
return result.get('embedding', [])
response = requests.post(url, json=payload)
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:
logging.error(f"Gemini embed_content error: {e}")
logging.error(f"Gemini embed_content parse error: {e}\n{traceback.format_exc()}")
return []
def create_embedding_adapter(
@@ -184,7 +206,6 @@ def create_embedding_adapter(
elif fmt == "ml studio":
return MLStudioEmbeddingAdapter(api_key, base_url, model_name)
elif fmt == "gemini":
# base_url 对 Gemini 暂无用处,可忽略
return GeminiEmbeddingAdapter(api_key, model_name)
return GeminiEmbeddingAdapter(api_key, model_name, base_url)
else:
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")
elif new_value == "OpenAI":
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":
self.embedding_url_var.set("https://[az].openai.azure.com/openai/deployments/[model]/embeddings?api-version=2023-05-15")
elif new_value == "DeepSeek":
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):
self.embeddings_config_tab.grid_rowconfigure(i, weight=0)
@@ -580,7 +584,7 @@ class NovelGeneratorGUI:
column=0,
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(
self.embeddings_config_tab,
values=emb_interface_options,