ForcePilot/src/models/embedding.py
Wenjie Zhang 782eb4ae31 refactor(chat): 优化异步调用及增强内容审查机制
- 将chat_router中的predict_async重命名为call_async并调用model.call替代model.predict
- 在流式消息处理中新增加基于LLM的内容审查
- 配置类中增加enable_content_guard_llm及对应LLM模型配置项
- 兼容models.private.yaml替代旧的models.private.yml文件名
- chat_model及embedding模块统一将predict方法重命名为call或encode,增强接口语义
- ContentGuard新增基于LLM的内容合规检测功能,支持动态加载审查模型
- 更新静态模型配置提示,建议使用models.private.yaml文件
- 新增示例CSV测试数据文件,补充测试用例基础数据
2025-09-19 00:57:53 +08:00

187 lines
7.6 KiB
Python
Raw 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.

import asyncio
import json
import os
from abc import ABC, abstractmethod
import httpx
import requests
from src import config
from src.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):
"""
Args:
model: 模型名称冗余设计同name
name: 模型名称冗余设计同model
dimension: 维度
url: 请求URL冗余设计同base_url
base_url: 基础URL请求URL冗余设计同url
api_key: 请求API密钥
"""
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.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")
async def aencode_queries(self, queries: list[str] | str) -> list[list[float]]:
"""等同于aencode"""
return await self.aencode(queries)
def batch_encode(self, messages: list[str], batch_size: int = 40) -> list[list[float]]:
# logger.info(f"Batch encoding {len(messages)} messages")
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})")
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 = 40) -> list[list[float]]:
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}
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
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]]:
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:
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:
logger.error(f"Ollama Embedding async request failed: {e}, {payload}")
raise ValueError(f"Ollama Embedding async request failed: {e}")
class OtherEmbedding(BaseEmbeddingModel):
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.headers = {"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"}
def build_payload(self, message: list[str] | str) -> dict:
return {"model": self.model, "input": message}
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:
logger.error(f"Other Embedding async request failed: {e}, {payload}")
raise ValueError(f"Other Embedding async request failed: {e}")
def select_embedding_model(model_id):
provider, model_name = model_id.split("/", 1) if model_id else ("", "")
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}"
logger.debug(f"Loading embedding model {model_id}")
if provider == "local":
raise ValueError("Local embedding model is not supported, please use other embedding models")
elif provider == "ollama":
model = OllamaEmbedding(**config.embed_model_names[model_id])
else:
model = OtherEmbedding(**config.embed_model_names[model_id])
return model