2024-07-21 18:15:28 +08:00
|
|
|
|
import os
|
2025-02-23 16:38:56 +08:00
|
|
|
|
import json
|
|
|
|
|
|
import requests
|
2025-04-23 21:18:39 +08:00
|
|
|
|
import asyncio
|
2025-05-09 23:45:16 +08:00
|
|
|
|
from abc import abstractmethod
|
2025-03-03 21:57:27 +08:00
|
|
|
|
from zhipuai import ZhipuAI
|
2025-05-09 23:45:16 +08:00
|
|
|
|
from langchain_huggingface import HuggingFaceEmbeddings
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
2025-03-20 19:51:46 +08:00
|
|
|
|
from src import config
|
2025-03-04 13:49:00 +08:00
|
|
|
|
from src.utils import hashstr, logger, get_docker_safe_url
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
|
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
class BaseEmbeddingModel:
|
2025-03-03 21:57:27 +08:00
|
|
|
|
embed_state = {}
|
2025-03-20 19:51:46 +08:00
|
|
|
|
|
2025-05-09 23:45:16 +08:00
|
|
|
|
@abstractmethod
|
|
|
|
|
|
def predict(self, message):
|
|
|
|
|
|
raise NotImplementedError("Subclasses must implement this method")
|
|
|
|
|
|
|
2025-03-20 19:51:46 +08:00
|
|
|
|
def get_dimension(self):
|
|
|
|
|
|
if hasattr(self, "dimension"):
|
|
|
|
|
|
return self.dimension
|
|
|
|
|
|
|
2025-04-06 20:47:45 +08:00
|
|
|
|
if hasattr(self, "embed_model_fullname"):
|
2025-04-06 20:33:15 +08:00
|
|
|
|
return config.embed_model_names[self.embed_model_fullname].get("dimension", None)
|
2025-03-20 19:51:46 +08:00
|
|
|
|
|
2025-04-06 20:33:15 +08:00
|
|
|
|
return config.embed_model_names[self.model].get("dimension", None)
|
2025-03-04 13:49:00 +08:00
|
|
|
|
|
|
|
|
|
|
def encode(self, message):
|
|
|
|
|
|
return self.predict(message)
|
|
|
|
|
|
|
|
|
|
|
|
def encode_queries(self, queries):
|
|
|
|
|
|
return self.predict(queries)
|
2025-03-03 21:57:27 +08:00
|
|
|
|
|
2025-04-23 21:18:39 +08:00
|
|
|
|
async def aencode(self, message):
|
|
|
|
|
|
return await asyncio.to_thread(self.encode, message)
|
|
|
|
|
|
|
|
|
|
|
|
async def aencode_queries(self, queries):
|
|
|
|
|
|
return await asyncio.to_thread(self.encode_queries, queries)
|
|
|
|
|
|
|
|
|
|
|
|
async def abatch_encode(self, messages, batch_size=20):
|
|
|
|
|
|
return await asyncio.to_thread(self.batch_encode, messages, batch_size)
|
|
|
|
|
|
|
2025-03-03 21:57:27 +08:00
|
|
|
|
def batch_encode(self, messages, batch_size=20):
|
2025-03-07 01:05:50 +08:00
|
|
|
|
logger.info(f"Batch encoding {len(messages)} messages")
|
2025-03-03 21:57:27 +08:00
|
|
|
|
data = []
|
|
|
|
|
|
|
|
|
|
|
|
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} to {i+batch_size} with {len(messages)} messages")
|
2025-03-11 16:26:55 +08:00
|
|
|
|
response = self.encode(group_msg)
|
2025-05-09 23:45:16 +08:00
|
|
|
|
# logger.debug(f"Response: {len(response)=}, {len(group_msg)=}, {len(response[0])=}")
|
2025-03-03 21:57:27 +08:00
|
|
|
|
data.extend(response)
|
|
|
|
|
|
|
|
|
|
|
|
if len(messages) > batch_size:
|
|
|
|
|
|
self.embed_state[task_id]['progress'] = len(messages)
|
|
|
|
|
|
self.embed_state[task_id]['status'] = 'completed'
|
|
|
|
|
|
|
|
|
|
|
|
return data
|
|
|
|
|
|
|
2025-05-09 23:45:16 +08:00
|
|
|
|
class LocalEmbeddingModel(BaseEmbeddingModel):
|
2025-05-24 11:29:45 +08:00
|
|
|
|
def __init__(self, **kwargs):
|
2025-04-06 20:33:15 +08:00
|
|
|
|
info = config.embed_model_names[config.embed_model]
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
self.model = config.model_local_paths.get(info["name"], info.get("local_path"))
|
|
|
|
|
|
self.model = self.model or info["name"]
|
2025-03-20 19:51:46 +08:00
|
|
|
|
self.dimension = info["dimension"]
|
|
|
|
|
|
self.embed_model_fullname = config.embed_model
|
2025-03-04 13:49:00 +08:00
|
|
|
|
|
2025-04-13 23:50:19 +08:00
|
|
|
|
if os.getenv("MODEL_DIR"):
|
|
|
|
|
|
if os.path.exists(_path := os.path.join(os.getenv("MODEL_DIR"), self.model)):
|
|
|
|
|
|
self.model = _path
|
|
|
|
|
|
else:
|
|
|
|
|
|
logger.warning(f"Local model `{info['name']}` not found in `{self.model}`, using `{info['name']}`")
|
|
|
|
|
|
|
2025-05-23 15:30:14 +08:00
|
|
|
|
logger.info(f"Loading local model `{info['name']}` from `{self.model}` with device `{config.device}`")
|
|
|
|
|
|
logger.debug("如果没配置任何路径的话,正常情况下会自动从 Huggingface 下载模型,如果遇到下载失败,可以尝试使用 HF_MIRROR 环境变量;"
|
2025-05-09 23:45:16 +08:00
|
|
|
|
f"如果还是不行,建议手动下载到某个文件夹,比如 {os.getenv('MODEL_DIR', '/models')}/BAAI/bge-m3 目录下;")
|
|
|
|
|
|
|
|
|
|
|
|
self.model = HuggingFaceEmbeddings(
|
|
|
|
|
|
model_name=self.model,
|
|
|
|
|
|
model_kwargs={'device': config.device},
|
|
|
|
|
|
encode_kwargs={
|
|
|
|
|
|
'normalize_embeddings': True,
|
|
|
|
|
|
'prompt_name': info.get("query_instruction", None),
|
|
|
|
|
|
},
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
logger.info(f"Embedding model {info['name']} loaded, {self.model=}")
|
|
|
|
|
|
|
|
|
|
|
|
def predict(self, message):
|
|
|
|
|
|
return self.model.embed_documents(message)
|
2025-03-04 13:49:00 +08:00
|
|
|
|
|
2025-05-09 23:45:16 +08:00
|
|
|
|
async def aencode(self, message):
|
|
|
|
|
|
return await self.model.aembed_documents(message)
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
2025-05-09 23:45:16 +08:00
|
|
|
|
def encode_queries(self, queries):
|
2025-05-24 11:29:45 +08:00
|
|
|
|
logger.warning("Huggingface Model 不支持批量 encode queries,因此使用训练实现")
|
2025-05-09 23:45:16 +08:00
|
|
|
|
data = []
|
|
|
|
|
|
for q in queries:
|
|
|
|
|
|
data.append(self.predict(q))
|
|
|
|
|
|
|
|
|
|
|
|
return data
|
2024-07-17 18:52:20 +08:00
|
|
|
|
|
2024-07-21 18:15:28 +08:00
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
class ZhipuEmbedding(BaseEmbeddingModel):
|
2024-08-25 12:34:35 +08:00
|
|
|
|
|
2025-05-24 11:29:45 +08:00
|
|
|
|
def __init__(self) -> None:
|
2025-02-23 16:38:56 +08:00
|
|
|
|
self.config = config
|
2025-04-06 20:33:15 +08:00
|
|
|
|
self.model = config.embed_model_names[config.embed_model]["name"]
|
|
|
|
|
|
self.dimension = config.embed_model_names[config.embed_model]["dimension"]
|
2025-02-23 16:38:56 +08:00
|
|
|
|
self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY"))
|
2025-03-20 19:51:46 +08:00
|
|
|
|
self.embed_model_fullname = config.embed_model
|
2024-09-09 17:07:03 +08:00
|
|
|
|
|
2025-02-23 16:38:56 +08:00
|
|
|
|
def predict(self, message):
|
|
|
|
|
|
response = self.client.embeddings.create(
|
|
|
|
|
|
model=self.model,
|
|
|
|
|
|
input=message,
|
|
|
|
|
|
)
|
|
|
|
|
|
data = [a.embedding for a in response.data]
|
2024-08-25 12:34:35 +08:00
|
|
|
|
return data
|
2024-07-22 00:00:54 +08:00
|
|
|
|
|
|
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
class OllamaEmbedding(BaseEmbeddingModel):
|
2025-05-24 11:29:45 +08:00
|
|
|
|
def __init__(self) -> None:
|
2025-04-06 20:33:15 +08:00
|
|
|
|
self.info = config.embed_model_names[config.embed_model]
|
2025-03-04 13:49:00 +08:00
|
|
|
|
self.model = self.info["name"]
|
2025-06-27 12:22:35 +08:00
|
|
|
|
self.url = self.info.get("base_url", "http://localhost:11434/api/embed")
|
2025-03-04 13:49:00 +08:00
|
|
|
|
self.url = get_docker_safe_url(self.url)
|
2025-03-20 19:51:46 +08:00
|
|
|
|
self.dimension = self.info.get("dimension", None)
|
|
|
|
|
|
self.embed_model_fullname = config.embed_model
|
2024-07-22 00:00:54 +08:00
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
def predict(self, message: list[str] | str):
|
|
|
|
|
|
if isinstance(message, str):
|
|
|
|
|
|
message = [message]
|
2024-07-22 00:00:54 +08:00
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
payload = {
|
|
|
|
|
|
"model": self.model,
|
|
|
|
|
|
"input": message,
|
|
|
|
|
|
}
|
|
|
|
|
|
response = requests.request("POST", self.url, json=payload)
|
|
|
|
|
|
response = json.loads(response.text)
|
|
|
|
|
|
assert response.get("embeddings"), f"Ollama Embedding failed: {response}"
|
|
|
|
|
|
return response["embeddings"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class OtherEmbedding(BaseEmbeddingModel):
|
2025-02-23 16:38:56 +08:00
|
|
|
|
|
2025-05-24 11:29:45 +08:00
|
|
|
|
def __init__(self) -> None:
|
2025-04-06 20:33:15 +08:00
|
|
|
|
self.info = config.embed_model_names[config.embed_model]
|
2025-03-20 19:51:46 +08:00
|
|
|
|
self.embed_model_fullname = config.embed_model
|
|
|
|
|
|
self.dimension = self.info.get("dimension", None)
|
2025-03-04 13:49:00 +08:00
|
|
|
|
self.model = self.info["name"]
|
2025-06-25 02:45:57 +08:00
|
|
|
|
self.api_key = os.getenv(self.info["api_key"], self.info["api_key"])
|
2025-06-27 12:22:35 +08:00
|
|
|
|
self.url = get_docker_safe_url(self.info["base_url"])
|
2025-03-04 13:49:00 +08:00
|
|
|
|
assert self.url and self.model, f"URL and model are required. Cur embed model: {config.embed_model}"
|
2025-02-23 16:38:56 +08:00
|
|
|
|
self.headers = {
|
2025-03-04 13:49:00 +08:00
|
|
|
|
"Authorization": f"Bearer {self.api_key}",
|
2025-02-23 16:38:56 +08:00
|
|
|
|
"Content-Type": "application/json"
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
def predict(self, message):
|
2025-02-23 16:38:56 +08:00
|
|
|
|
payload = self.build_payload(message)
|
|
|
|
|
|
response = requests.request("POST", self.url, json=payload, headers=self.headers)
|
|
|
|
|
|
response = json.loads(response.text)
|
2025-03-04 13:49:00 +08:00
|
|
|
|
assert response["data"], f"Other Embedding failed: {response}"
|
2025-02-23 16:38:56 +08:00
|
|
|
|
data = [a["embedding"] for a in response["data"]]
|
|
|
|
|
|
return data
|
|
|
|
|
|
|
|
|
|
|
|
def build_payload(self, message):
|
|
|
|
|
|
return {
|
|
|
|
|
|
"model": self.model,
|
|
|
|
|
|
"input": message,
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2025-05-24 11:29:45 +08:00
|
|
|
|
def get_embedding_model():
|
2025-02-23 16:38:56 +08:00
|
|
|
|
provider, model_name = config.embed_model.split('/', 1)
|
2025-05-24 11:29:45 +08:00
|
|
|
|
support_embed_models = config.embed_model_names.keys()
|
|
|
|
|
|
assert config.embed_model in support_embed_models, f"Unsupported embed model: {config.embed_model}, only support {support_embed_models}"
|
2025-02-23 16:38:56 +08:00
|
|
|
|
logger.debug(f"Loading embedding model {config.embed_model}")
|
|
|
|
|
|
if provider == "local":
|
2025-06-25 02:45:57 +08:00
|
|
|
|
logger.warning("[DEPRECATED] Local embedding model will be removed in v0.2, please use other embedding models")
|
2025-05-24 11:29:45 +08:00
|
|
|
|
model = LocalEmbeddingModel()
|
2024-08-25 20:29:24 +08:00
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
elif provider == "zhipu":
|
2025-05-24 11:29:45 +08:00
|
|
|
|
model = ZhipuEmbedding()
|
2024-08-25 20:29:24 +08:00
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
elif provider == "ollama":
|
2025-05-24 11:29:45 +08:00
|
|
|
|
model = OllamaEmbedding()
|
2025-03-04 13:49:00 +08:00
|
|
|
|
|
|
|
|
|
|
else:
|
2025-05-24 11:29:45 +08:00
|
|
|
|
model = OtherEmbedding()
|
2024-08-25 20:29:24 +08:00
|
|
|
|
|
2024-09-11 01:07:19 +08:00
|
|
|
|
return model
|
|
|
|
|
|
|
|
|
|
|
|
def handle_local_model(paths, model_name, default_path):
|
|
|
|
|
|
model_path = paths.get(model_name, default_path)
|
2025-05-24 11:29:45 +08:00
|
|
|
|
return model_path
|