207 lines
7.1 KiB
Python
207 lines
7.1 KiB
Python
"""
|
||
Embedding 服务
|
||
|
||
通过 OpenAI 兼容的 Embeddings API 将文本转换为向量
|
||
支持所有兼容 OpenAI 接口的提供商(OpenAI、Qwen、Ollama 等)
|
||
"""
|
||
import logging
|
||
import math
|
||
from typing import List, Optional
|
||
|
||
from sqlalchemy import select
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
# 单次批量请求的默认最大文本数
|
||
DEFAULT_BATCH_SIZE = 50
|
||
|
||
# 不同提供商的批次大小限制
|
||
PROVIDER_BATCH_SIZE = {
|
||
'qwen': 10, # 阿里云 DashScope 限制 10
|
||
'dashscope': 10,
|
||
'siliconflow': 10, # 硅基流动限制较小
|
||
'ollama': 1, # Ollama 通常逐条处理
|
||
}
|
||
|
||
|
||
|
||
class EmbeddingService:
|
||
"""
|
||
Embedding 服务
|
||
|
||
通过模型 ID 获取对应的提供商,调用 OpenAI 兼容的 Embeddings API
|
||
"""
|
||
|
||
def __init__(self, db: AsyncSession):
|
||
self._db = db
|
||
self._client_cache = {}
|
||
|
||
async def _get_client_and_model(self, model_id: str):
|
||
"""
|
||
根据模型 ID 获取异步客户端和模型名称
|
||
|
||
Returns:
|
||
(async_client, model_name, max_tokens, provider_type)
|
||
"""
|
||
from ai_platform.models import LLMModel, LLMProvider
|
||
|
||
result = await self._db.execute(
|
||
select(LLMModel).where(
|
||
LLMModel.id == model_id,
|
||
LLMModel.is_active == True,
|
||
LLMModel.is_deleted == False
|
||
)
|
||
)
|
||
model = result.scalar_one_or_none()
|
||
if not model:
|
||
raise ValueError(f'Embedding 模型不存在或已禁用: {model_id}')
|
||
if model.model_type != 'embedding':
|
||
raise ValueError(f'模型 {model.display_name} 不是 Embedding 类型')
|
||
|
||
provider_result = await self._db.execute(
|
||
select(LLMProvider).where(
|
||
LLMProvider.id == model.provider_id,
|
||
LLMProvider.is_active == True,
|
||
LLMProvider.is_deleted == False
|
||
)
|
||
)
|
||
provider = provider_result.scalar_one_or_none()
|
||
if not provider:
|
||
raise ValueError('Embedding 模型对应的提供商不存在或已禁用')
|
||
|
||
cache_key = str(provider.id)
|
||
if cache_key not in self._client_cache:
|
||
import httpx
|
||
from openai import AsyncOpenAI
|
||
|
||
# 根据提供商类型确定 base_url
|
||
if provider.provider_type == 'ollama':
|
||
base_url = (provider.ollama_host or 'http://localhost:11434').rstrip('/') + '/v1'
|
||
else:
|
||
from ai_platform.providers.registry import ProviderRegistry
|
||
provider_class = ProviderRegistry.get(provider.provider_type)
|
||
default_base = getattr(provider_class, 'DEFAULT_API_BASE', 'https://api.openai.com/v1') if provider_class else 'https://api.openai.com/v1'
|
||
base_url = provider.api_base or default_base
|
||
|
||
self._client_cache[cache_key] = AsyncOpenAI(
|
||
api_key=provider.api_key or 'ollama',
|
||
base_url=base_url,
|
||
timeout=httpx.Timeout(120.0, connect=30.0),
|
||
max_retries=5,
|
||
)
|
||
|
||
return self._client_cache[cache_key], model.model_name, model.context_window or 8191, provider.provider_type
|
||
|
||
async def embed_text(
|
||
self,
|
||
model_id: str,
|
||
text: str,
|
||
dimensions: Optional[int] = None,
|
||
) -> List[float]:
|
||
"""
|
||
将单个文本转换为向量
|
||
|
||
Args:
|
||
model_id: Embedding 模型 ID
|
||
text: 文本内容
|
||
dimensions: 向量维度(可选,部分模型支持)
|
||
|
||
Returns:
|
||
向量列表 List[float]
|
||
"""
|
||
results = await self.embed_texts(model_id, [text], dimensions)
|
||
return results[0]
|
||
|
||
async def embed_texts(
|
||
self,
|
||
model_id: str,
|
||
texts: List[str],
|
||
dimensions: Optional[int] = None,
|
||
) -> List[List[float]]:
|
||
"""
|
||
批量将文本转换为向量
|
||
|
||
Args:
|
||
model_id: Embedding 模型 ID
|
||
texts: 文本列表
|
||
dimensions: 向量维度(可选)
|
||
|
||
Returns:
|
||
向量列表 List[List[float]]
|
||
"""
|
||
if not texts:
|
||
return []
|
||
|
||
client, model_name, max_tokens, provider_type = await self._get_client_and_model(model_id)
|
||
|
||
# 根据提供商类型确定批次大小
|
||
batch_size = PROVIDER_BATCH_SIZE.get(provider_type, DEFAULT_BATCH_SIZE)
|
||
|
||
# 预处理:截断超长文本
|
||
processed_texts = []
|
||
for text in texts:
|
||
if not text or not text.strip():
|
||
processed_texts.append(" ")
|
||
else:
|
||
# 粗略估算 token 数(中文约 1 字 = 1.5 token,英文约 4 字符 = 1 token)
|
||
# 保守截断到 max_tokens * 2 个字符
|
||
max_chars = max_tokens * 2
|
||
if len(text) > max_chars:
|
||
processed_texts.append(text[:max_chars])
|
||
else:
|
||
processed_texts.append(text)
|
||
|
||
# 分批处理
|
||
all_embeddings = [None] * len(processed_texts)
|
||
total_batches = math.ceil(len(processed_texts) / batch_size)
|
||
|
||
for batch_idx in range(total_batches):
|
||
start = batch_idx * batch_size
|
||
end = min(start + batch_size, len(processed_texts))
|
||
batch_texts = processed_texts[start:end]
|
||
|
||
try:
|
||
kwargs = {
|
||
'model': model_name,
|
||
'input': batch_texts,
|
||
}
|
||
# 仅对明确支持 dimensions 参数的模型传递该参数
|
||
if dimensions:
|
||
model_lower = model_name.lower()
|
||
# OpenAI text-embedding-3 系列原生支持任意 dimensions
|
||
if 'text-embedding-3' in model_lower:
|
||
kwargs['dimensions'] = dimensions
|
||
# DashScope text-embedding-v3 只接受 [64,128,256,512,768,1024]
|
||
elif 'text-embedding-v3' in model_lower and dimensions in (64, 128, 256, 512, 768, 1024):
|
||
kwargs['dimensions'] = dimensions
|
||
|
||
response = await client.embeddings.create(**kwargs)
|
||
|
||
for item in response.data:
|
||
all_embeddings[start + item.index] = item.embedding
|
||
|
||
except Exception as e:
|
||
logger.exception(f'Embedding 批次 {batch_idx + 1}/{total_batches} 失败: {e}')
|
||
raise ValueError(f'Embedding 调用失败: {str(e)}')
|
||
|
||
# 检查是否所有向量都已生成
|
||
for i, emb in enumerate(all_embeddings):
|
||
if emb is None:
|
||
raise ValueError(f'第 {i} 个文本的向量未生成')
|
||
|
||
return all_embeddings
|
||
|
||
async def get_embedding_dimensions(self, model_id: str) -> int:
|
||
"""
|
||
获取模型的向量维度(通过嵌入一个测试文本来检测)
|
||
|
||
Args:
|
||
model_id: Embedding 模型 ID
|
||
|
||
Returns:
|
||
向量维度
|
||
"""
|
||
test_embedding = await self.embed_text(model_id, "test")
|
||
return len(test_embedding)
|