ForcePilot/src/models/embedding.py

117 lines
4.0 KiB
Python
Raw Normal View History

2024-07-21 18:15:28 +08:00
import os
2024-09-09 17:07:03 +08:00
import uuid
2024-07-17 18:52:20 +08:00
from FlagEmbedding import FlagModel, FlagReranker
2024-08-25 20:29:24 +08:00
from src.config import EMBED_MODEL_INFO, RERANKER_LIST
from src.utils.logging_config import setup_logger
2024-09-09 17:07:03 +08:00
from src.utils import hashstr
logger = setup_logger("EmbeddingModel")
2024-09-09 17:07:03 +08:00
GLOBAL_EMBED_STATE = {}
class EmbeddingModel(FlagModel):
2024-08-25 20:29:24 +08:00
def __init__(self, model_info, config, **kwargs):
self.info = model_info
2024-09-11 01:07:19 +08:00
model_name_or_path = handle_local_model(
paths=config.model_local_paths,
model_name=model_info.name,
default_path=model_info.default_path)
2024-08-25 20:29:24 +08:00
logger.info(f"Loading embedding model {model_info.name} from {model_name_or_path}")
super().__init__(model_name_or_path,
2024-08-25 20:29:24 +08:00
query_instruction_for_retrieval=model_info.get("query_instruction", None),
use_fp16=False, **kwargs)
2024-08-25 20:29:24 +08:00
logger.info(f"Embedding model {model_info.name} loaded")
2024-07-17 18:52:20 +08:00
2024-07-21 18:15:28 +08:00
class Reranker(FlagReranker):
2024-07-17 18:52:20 +08:00
def __init__(self, config, **kwargs):
2024-07-21 18:15:28 +08:00
assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}"
2024-07-17 18:52:20 +08:00
2024-09-11 01:07:19 +08:00
model_name_or_path = handle_local_model(
paths=config.model_local_paths,
model_name=config.reranker,
default_path=RERANKER_LIST[config.reranker])
logger.info(f"Loading Reranker model {config.reranker} from {model_name_or_path}")
2024-07-17 18:52:20 +08:00
super().__init__(model_name_or_path, use_fp16=True, **kwargs)
logger.info(f"Reranker model {config.reranker} loaded")
2024-07-21 18:15:28 +08:00
from zhipuai import ZhipuAI
class ZhipuEmbedding:
2024-08-25 20:29:24 +08:00
def __init__(self, model_info, config) -> None:
2024-07-21 18:15:28 +08:00
self.config = config
2024-08-25 20:29:24 +08:00
self.model_info = model_info
2024-09-28 00:39:31 +08:00
self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY"))
logger.info("Zhipu Embedding model loaded")
self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:"
2024-07-21 18:15:28 +08:00
def predict(self, message):
2024-08-25 12:34:35 +08:00
data = []
2024-09-09 17:07:03 +08:00
if len(message) > 10:
global GLOBAL_EMBED_STATE
task_id = hashstr(message)
logger.info(f"Creating new state for process {task_id}")
GLOBAL_EMBED_STATE[task_id] = {
'status': 'in-progress',
'total': len(message),
'progress': 0
}
2024-08-25 12:34:35 +08:00
for i in range(0, len(message), 10):
2024-09-06 12:54:17 +08:00
if len(message) > 10:
logger.info(f"Encoding {i} to {i+10} with {len(message)} messages")
2024-09-09 17:07:03 +08:00
GLOBAL_EMBED_STATE[task_id]['progress'] = i
2024-08-25 12:34:35 +08:00
group_msg = message[i:i+10]
response = self.client.embeddings.create(
2024-08-25 20:29:24 +08:00
model=self.model_info.default_path,
input=group_msg,
2024-08-25 12:34:35 +08:00
)
data.extend([a.embedding for a in response.data])
2024-09-09 17:07:03 +08:00
if len(message) > 10:
GLOBAL_EMBED_STATE[task_id]['progress'] = len(message)
GLOBAL_EMBED_STATE[task_id]['status'] = 'completed'
2024-08-25 12:34:35 +08:00
return data
def encode(self, message):
return self.predict(message)
def encode_queries(self, queries):
# queries = [self.query_instruction_for_retrieval + query for query in queries]
return self.predict(queries)
def get_embedding_model(config):
2024-07-31 20:22:05 +08:00
if not config.enable_knowledge_base:
return None
2024-08-25 20:29:24 +08:00
assert config.embed_model in EMBED_MODEL_INFO.keys(), f"Unsupported embed model: {config.embed_model}, only support {EMBED_MODEL_INFO.keys()}"
if config.embed_model in ["bge-large-zh-v1.5"]:
model = EmbeddingModel(EMBED_MODEL_INFO[config.embed_model], config)
if config.embed_model in ["zhipu-embedding-2", "zhipu-embedding-3"]:
model = ZhipuEmbedding(EMBED_MODEL_INFO[config.embed_model], config)
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)
if os.getenv("MODEL_ROOT_DIR") and not os.path.isabs(model_path):
model_path = os.path.join(os.getenv("MODEL_ROOT_DIR"), model_path)
return model_path