Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,206 @@
|
||||
"""
|
||||
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)
|
||||
Reference in New Issue
Block a user