- 将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测试数据文件,补充测试用例基础数据
187 lines
7.6 KiB
Python
187 lines
7.6 KiB
Python
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
|