Files
2026-06-08 18:14:59 +08:00

207 lines
7.1 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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)