This commit is contained in:
YILING0013
2025-02-11 21:32:05 +08:00
parent d54871b2c2
commit 8ed0166d3e
2 changed files with 51 additions and 14 deletions
+41 -5
View File
@@ -45,7 +45,7 @@ class OpenAIEmbeddingAdapter(BaseEmbeddingAdapter):
def embed_query(self, query: str) -> List[float]: def embed_query(self, query: str) -> List[float]:
return self._embedding.embed_query(query) return self._embedding.embed_query(query)
class AzureOpenAIEmbeddingAdapter(BaseEmbeddingAdapter): class AzureOpenAIEmbeddingAdapter(BaseEmbeddingAdapter):
""" """
基于 AzureOpenAIEmbeddings(或兼容接口)的适配器 基于 AzureOpenAIEmbeddings(或兼容接口)的适配器
@@ -133,6 +133,38 @@ class MLStudioEmbeddingAdapter(BaseEmbeddingAdapter):
def embed_query(self, query: str) -> List[float]: def embed_query(self, query: str) -> List[float]:
return self._embedding.embed_query(query) return self._embedding.embed_query(query)
class GeminiEmbeddingAdapter(BaseEmbeddingAdapter):
"""
基于 Google Generative AI Gemini)接口的 Embedding 适配器
"""
def __init__(self, api_key: str, model_name: str):
from google import genai
# 全局配置,也可根据需要改成 Client(...) 初始化方式
genai.configure(api_key=api_key)
self.model_name = model_name
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([])
return embeddings
def embed_query(self, query: str) -> List[float]:
from google import genai
try:
result = genai.embed_content(model=self.model_name, content=query)
return result.get('embedding', [])
except Exception as e:
logging.error(f"Gemini embed_content error: {e}")
return []
def create_embedding_adapter( def create_embedding_adapter(
interface_format: str, interface_format: str,
api_key: str, api_key: str,
@@ -142,13 +174,17 @@ def create_embedding_adapter(
""" """
工厂函数:根据 interface_format 返回不同的 embedding 适配器实例 工厂函数:根据 interface_format 返回不同的 embedding 适配器实例
""" """
if interface_format.lower() == "openai": fmt = interface_format.strip().lower()
if fmt == "openai":
return OpenAIEmbeddingAdapter(api_key, base_url, model_name) return OpenAIEmbeddingAdapter(api_key, base_url, model_name)
elif interface_format.lower() == "azure openai": elif fmt == "azure openai":
return AzureOpenAIEmbeddingAdapter(api_key, base_url, model_name) return AzureOpenAIEmbeddingAdapter(api_key, base_url, model_name)
elif interface_format.lower() == "ollama": elif fmt == "ollama":
return OllamaEmbeddingAdapter(model_name, base_url) return OllamaEmbeddingAdapter(model_name, base_url)
elif interface_format.lower() == "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":
# base_url 对 Gemini 暂无用处,可忽略
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}")
+10 -9
View File
@@ -18,7 +18,7 @@ def ensure_openai_base_url_has_v1(url: str) -> str:
class BaseLLMAdapter: class BaseLLMAdapter:
""" """
统一的 LLM 接口基类,为不同后端(OpenAI、Ollama、ML Studio 等)提供一致的方法签名。 统一的 LLM 接口基类,为不同后端(OpenAI、Ollama、ML Studio、Gemini等)提供一致的方法签名。
""" """
def invoke(self, prompt: str) -> str: def invoke(self, prompt: str) -> str:
raise NotImplementedError("Subclasses must implement .invoke(prompt) method.") raise NotImplementedError("Subclasses must implement .invoke(prompt) method.")
@@ -81,7 +81,7 @@ class OpenAIAdapter(BaseLLMAdapter):
class GeminiAdapter(BaseLLMAdapter): class GeminiAdapter(BaseLLMAdapter):
""" """
适配 Google Gemini 接口 适配 Google Gemini (Google Generative AI) 接口
""" """
def __init__(self, api_key: str, model_name: str, max_tokens: int, temperature: float = 0.7, timeout: Optional[int] = 600): def __init__(self, api_key: str, model_name: str, max_tokens: int, temperature: float = 0.7, timeout: Optional[int] = 600):
self.api_key = api_key self.api_key = api_key
@@ -151,7 +151,6 @@ class AzureOpenAIAdapter(BaseLLMAdapter):
class OllamaAdapter(BaseLLMAdapter): class OllamaAdapter(BaseLLMAdapter):
""" """
Ollama 同样有一个 OpenAI-like /v1/chat 接口,可直接使用 ChatOpenAI。 Ollama 同样有一个 OpenAI-like /v1/chat 接口,可直接使用 ChatOpenAI。
但是通常 Ollama 默认本地服务在 http://localhost:11434,如果符合OpenAI风格即可直接传参。
""" """
def __init__(self, api_key: str, base_url: str, model_name: str, max_tokens: int, temperature: float = 0.7, timeout: Optional[int] = 600): def __init__(self, api_key: str, base_url: str, model_name: str, max_tokens: int, temperature: float = 0.7, timeout: Optional[int] = 600):
self.base_url = ensure_openai_base_url_has_v1(base_url) self.base_url = ensure_openai_base_url_has_v1(base_url)
@@ -214,17 +213,19 @@ def create_llm_adapter(
""" """
工厂函数:根据 interface_format 返回不同的适配器实例。 工厂函数:根据 interface_format 返回不同的适配器实例。
""" """
if interface_format.lower() == "deepseek": fmt = interface_format.strip().lower()
if fmt == "deepseek":
return DeepSeekAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) return DeepSeekAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
elif interface_format.lower() == "openai": elif fmt == "openai":
return OpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) return OpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
elif interface_format.lower() == "azure openai": elif fmt == "azure openai":
return AzureOpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) return AzureOpenAIAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
elif interface_format.lower() == "ollama": elif fmt == "ollama":
return OllamaAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) return OllamaAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
elif interface_format.lower() == "ml studio": elif fmt == "ml studio":
return MLStudioAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout) return MLStudioAdapter(api_key, base_url, model_name, max_tokens, temperature, timeout)
elif interface_format.lower() == "gemini": elif fmt == "gemini":
# base_url 对 Gemini 暂无用处,可忽略
return GeminiAdapter(api_key, model_name, max_tokens, temperature, timeout) return GeminiAdapter(api_key, model_name, max_tokens, temperature, timeout)
else: else:
raise ValueError(f"Unknown interface_format: {interface_format}") raise ValueError(f"Unknown interface_format: {interface_format}")