ForcePilot/backend/package/yuxi/models/embed.py

408 lines
15 KiB
Python
Raw Normal View History

import asyncio
2025-02-23 16:38:56 +08:00
import json
import os
from abc import ABC, abstractmethod
import httpx
2025-02-23 16:38:56 +08:00
import requests
from yuxi import config
from yuxi.services.model_cache import is_v2_spec_format
from yuxi.utils import get_docker_safe_url, hashstr, logger
class BaseEmbeddingModel(ABC):
def __init__(
self,
model=None,
name=None,
dimension=None,
url=None,
base_url=None,
api_key=None,
model_id=None,
batch_size=40,
):
"""
Args:
model: 模型名称冗余设计同name
name: 模型名称冗余设计同model
dimension: 维度
url: 请求URL冗余设计同base_url
base_url: 基础URL请求URL冗余设计同url
api_key: 请求API密钥
batch_size: 模型推荐的批量向量化大小
"""
base_url = base_url or url
self.model = model or name
self.dimension = dimension
self.base_url = get_docker_safe_url(base_url)
self.api_key = os.getenv(api_key, api_key)
self.batch_size = int(batch_size or 40)
self.embed_state = {}
@abstractmethod
def encode(self, message: list[str] | str) -> list[list[float]]:
"""同步编码"""
raise NotImplementedError("Subclasses must implement this method")
def encode_queries(self, queries: list[str] | str) -> list[list[float]]:
"""等同于encode"""
return self.encode(queries)
@abstractmethod
async def aencode(self, message: list[str] | str) -> list[list[float]]:
"""异步编码"""
raise NotImplementedError("Subclasses must implement this method")
2025-04-23 21:18:39 +08:00
async def aencode_queries(self, queries: list[str] | str) -> list[list[float]]:
"""等同于aencode"""
return await self.aencode(queries)
2025-04-23 21:18:39 +08:00
def batch_encode(self, messages: list[str], batch_size: int | None = None) -> list[list[float]]:
# logger.info(f"Batch encoding {len(messages)} messages")
batch_size = batch_size or self.batch_size
data = []
task_id = None
if len(messages) > batch_size:
task_id = hashstr(messages)
self.embed_state[task_id] = {"status": "in-progress", "total": len(messages), "progress": 0}
for i in range(0, len(messages), batch_size):
group_msg = messages[i : i + batch_size]
logger.info(f"Encoding [{i}/{len(messages)}] messages (bsz={batch_size})")
2025-03-11 16:26:55 +08:00
response = self.encode(group_msg)
data.extend(response)
if task_id:
self.embed_state[task_id]["progress"] = i + len(group_msg)
if task_id:
self.embed_state[task_id]["status"] = "completed"
return data
async def abatch_encode(self, messages: list[str], batch_size: int | None = None) -> list[list[float]]:
batch_size = batch_size or self.batch_size
data = []
task_id = None
if len(messages) > batch_size:
task_id = hashstr(messages)
self.embed_state[task_id] = {"status": "in-progress", "total": len(messages), "progress": 0}
2026-02-28 16:33:54 +08:00
# 保留原有逻辑:
# 使用 asyncio.gather 并发执行所有 embedding 批次请求:
2026-02-26 12:43:38 +08:00
# tasks = []
# for i in range(0, len(messages), batch_size):
# group_msg = messages[i : i + batch_size]
# tasks.append(self.aencode(group_msg))
# results = await asyncio.gather(*tasks)
# for res in results:
# data.extend(res)
# if task_id:
# self.embed_state[task_id]["progress"] = len(messages)
# self.embed_state[task_id]["status"] = "completed"
# return data
2026-02-28 16:33:54 +08:00
for i in range(0, len(messages), batch_size):
group_msg = messages[i : i + batch_size]
2026-02-26 12:43:38 +08:00
logger.info(f"Async encoding [{i}/{len(messages)}] messages (bsz={batch_size})")
res = await self.aencode(group_msg)
data.extend(res)
2026-02-26 12:43:38 +08:00
if task_id:
self.embed_state[task_id]["progress"] = i + len(group_msg)
if task_id:
self.embed_state[task_id]["status"] = "completed"
return data
async def test_connection(self) -> tuple[bool, str]:
"""
测试embedding模型的连接性
Returns:
tuple: (success: bool, message: str)
"""
try:
# 使用简单的测试文本
test_text = ["Hello world"]
await self.aencode(test_text)
return True, "连接正常"
except Exception as e:
error_msg = str(e)
error_msg += f", maybe you can check the `{self.base_url}` end with /embeddings as examples."
logger.error(error_msg)
return False, error_msg
2025-03-04 13:49:00 +08:00
class OllamaEmbedding(BaseEmbeddingModel):
"""
Ollama Embedding Model
"""
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.base_url = self.base_url or get_docker_safe_url("http://localhost:11434/api/embed")
def encode(self, message: list[str] | str) -> list[list[float]]:
2025-03-04 13:49:00 +08:00
if isinstance(message, str):
message = [message]
payload = {"model": self.model, "input": message}
try:
response = requests.post(self.base_url, json=payload, timeout=60)
response.raise_for_status()
result = response.json()
if "embeddings" not in result:
raise ValueError(f"Ollama Embedding failed: Invalid response format {result}")
return result["embeddings"]
except (requests.RequestException, json.JSONDecodeError) as e:
logger.error(f"Ollama Embedding request failed: {e}, {payload}")
raise ValueError(f"Ollama Embedding request failed: {e}")
async def aencode(self, message: list[str] | str) -> list[list[float]]:
if isinstance(message, str):
message = [message]
payload = {"model": self.model, "input": message}
async with httpx.AsyncClient() as client:
try:
print(f"\n\n\nOllama Embedding request: {payload}\n\n\n")
response = await client.post(self.base_url, json=payload, timeout=60)
response.raise_for_status()
result = response.json()
if "embeddings" not in result:
raise ValueError(f"Ollama Embedding failed: Invalid response format {result}")
return result["embeddings"]
except (httpx.RequestError, json.JSONDecodeError) as e:
raise ValueError(f"Ollama Embedding async request failed: {e}, {payload}, {self.base_url=}")
2025-11-05 16:22:51 +08:00
2025-03-04 13:49:00 +08:00
class OtherEmbedding(BaseEmbeddingModel):
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.headers = {"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"}
2025-02-23 16:38:56 +08:00
def build_payload(self, message: list[str] | str) -> dict:
return {"model": self.model, "input": message}
2025-02-23 16:38:56 +08:00
def encode(self, message: list[str] | str) -> list[list[float]]:
payload = self.build_payload(message)
try:
response = requests.post(self.base_url, json=payload, headers=self.headers, timeout=60)
response.raise_for_status()
result = response.json()
if not isinstance(result, dict) or "data" not in result:
raise ValueError(f"Other Embedding failed: Invalid response format {result}")
return [item["embedding"] for item in result["data"]]
except (requests.RequestException, json.JSONDecodeError) as e:
logger.error(f"Other Embedding request failed: {e}, {payload}")
raise ValueError(f"Other Embedding request failed: {e}")
async def aencode(self, message: list[str] | str) -> list[list[float]]:
payload = self.build_payload(message)
async with httpx.AsyncClient() as client:
try:
response = await client.post(self.base_url, json=payload, headers=self.headers, timeout=60)
response.raise_for_status()
result = response.json()
if not isinstance(result, dict) or "data" not in result:
raise ValueError(f"Other Embedding failed: Invalid response format {result}")
return [item["embedding"] for item in result["data"]]
except (httpx.RequestError, json.JSONDecodeError) as e:
raise ValueError(f"Other Embedding async request failed: {e}, {payload}, {self.base_url=}")
async def test_embedding_model_status(model_id: str) -> dict:
"""
测试指定embedding模型的状态
Args:
model_id: 模型ID格式为 "provider/model_name"
Returns:
dict: 包含状态信息的字典
"""
try:
support_embed_models = config.embed_model_names.keys()
if model_id not in support_embed_models:
return {"model_id": model_id, "status": "unsupported", "message": f"不支持的模型: {model_id}"}
# 选择并创建模型实例
model = select_embedding_model(model_id)
# 测试连接
success, message = await model.test_connection()
return {
"model_id": model_id,
"status": "available" if success else "unavailable",
"message": message if not success else "连接正常",
"dimension": model.dimension,
}
except Exception as e:
logger.warning(f"测试embedding模型状态失败 {model_id}: {e}")
return {"model_id": model_id, "status": "error", "message": str(e)}
async def test_all_embedding_models_status() -> dict:
"""
测试所有支持的embedding模型状态
Returns:
dict: 包含所有模型状态的字典
"""
support_embed_models = list(config.embed_model_names.keys())
results = {}
# 并发测试所有模型
tasks = [test_embedding_model_status(model_id) for model_id in support_embed_models]
model_statuses = await asyncio.gather(*tasks, return_exceptions=True)
for i, status in enumerate(model_statuses):
if isinstance(status, Exception):
model_id = support_embed_models[i]
results[model_id] = {"model_id": model_id, "status": "error", "message": str(status)}
else:
results[status["model_id"]] = status
return {
"models": results,
"total": len(support_embed_models),
"available": len([m for m in results.values() if m["status"] == "available"]),
}
def get_embedding_model_info_by_id(model_id: str) -> dict:
"""
通过模型ID获取Embedding模型的标准化配置信息统一入口V1/V2 自动识别
V1 格式: provider/model_name "siliconflow/BAAI/bge-m3" config.embed_model_names 查找
V2 格式: provider_id:model_id "siliconflow:BAAI/bge-m3" model_cache 查找
Returns:
dict: 包含 base_urlapi_keynamedimensionmodel_idbatch_size 等字段的配置字典
api_key 已从环境变量解析为实际值
"""
# V2 spec 检测
if isinstance(model_id, str) and is_v2_spec_format(model_id):
from yuxi.services.model_cache import model_cache
info = model_cache.get_model_info(model_id)
if info:
logger.info(f"Loaded v2 embedding model info for {model_id}")
return {
"name": info.display_name,
"dimension": info.dimension,
"base_url": info.base_url,
"api_key": info.api_key,
"model_id": info.spec,
"batch_size": info.batch_size,
}
support_embed_models = config.embed_model_names.keys()
assert model_id in support_embed_models, f"Unsupported embed model: {model_id}, only support {support_embed_models}"
embed_config = config.embed_model_names[model_id].model_dump()
# 解析 api_key如果值本身是环境变量名则从环境变量获取实际值
embed_config["api_key"] = os.getenv(embed_config["api_key"]) or embed_config["api_key"]
logger.info(f"Loaded embedding model info for {model_id}")
return embed_config
def select_embedding_model(model_id):
"""选择 Embedding 模型V1/V2 自动识别)。
V1 格式: provider/model_name斜杠分隔
V2 格式: provider_id:model_id冒号分隔
"""
# V2 spec 检测:第一个特殊字符为冒号且存在于缓存中
if isinstance(model_id, str) and is_v2_spec_format(model_id):
from yuxi.services.model_cache import model_cache
if model_cache.is_v2_spec(model_id):
return select_embedding_model_v2(model_id)
provider, model_name = model_id.split("/", 1) if model_id else ("", "")
logger.info(f"Loading embedding model {model_id}")
if provider == "local":
raise ValueError("Local embedding model is not supported, please use other embedding models")
embed_config = get_embedding_model_info_by_id(model_id)
2025-10-22 11:51:32 +08:00
if provider == "ollama":
model = OllamaEmbedding(**embed_config)
else:
2025-10-22 11:51:32 +08:00
model = OtherEmbedding(**embed_config)
return model
def select_embedding_model_v2(spec: str):
"""根据 v2 specprovider_id:model_id选择 Embedding 模型。
v2 spec 格式使用冒号分隔: siliconflow:BAAI/bge-m3
数据来源为数据库中的 model_providers 通过全局缓存访问
"""
from yuxi.services.model_cache import model_cache
info = model_cache.get_model_info(spec)
if not info:
raise ValueError(f"Unknown v2 embedding model spec: {spec}")
if info.model_type != "embedding":
raise ValueError(f"Model {spec} is not an embedding model (type={info.model_type})")
logger.info(f"Selecting v2 embedding model: {spec} (provider_type={info.provider_type})")
if info.provider_type == "ollama":
return OllamaEmbedding(
model=info.model_id,
base_url=info.base_url,
api_key=info.api_key,
dimension=info.dimension,
batch_size=info.batch_size,
)
else:
return OtherEmbedding(
model=info.model_id,
base_url=info.base_url,
api_key=info.api_key,
dimension=info.dimension,
batch_size=info.batch_size,
)
async def test_embedding_model_status_by_spec(spec: str) -> dict:
"""根据 full spec 测试 Embedding 模型状态(自动识别 V1/V2
V1 spec 格式: provider/model_name斜杠分隔
V2 spec 格式: provider_id:model_id冒号分隔
"""
try:
if is_v2_spec_format(spec):
from yuxi.services.model_cache import model_cache
if model_cache.is_v2_spec(spec):
model = select_embedding_model_v2(spec)
else:
return {"spec": spec, "status": "unsupported", "message": f"不支持的 V2 模型: {spec}"}
else:
model = select_embedding_model(spec)
success, message = await model.test_connection()
return {
"spec": spec,
"status": "available" if success else "unavailable",
"message": "连接正常" if success else message,
}
except Exception as e:
logger.warning(f"测试 Embedding 模型状态失败 {spec}: {e}")
return {"spec": spec, "status": "error", "message": str(e)}