修复本地embedding模型的运行bug,添加了 batch_encode 的方法
This commit is contained in:
parent
17163a3b86
commit
e998102eeb
@ -2,28 +2,12 @@ import os
|
|||||||
import json
|
import json
|
||||||
import requests
|
import requests
|
||||||
from FlagEmbedding import FlagModel
|
from FlagEmbedding import FlagModel
|
||||||
|
from zhipuai import ZhipuAI
|
||||||
|
|
||||||
from src.config import EMBED_MODEL_INFO
|
from src.config import EMBED_MODEL_INFO
|
||||||
from src.utils import hashstr, logger
|
from src.utils import hashstr, logger
|
||||||
|
|
||||||
|
|
||||||
class LocalEmbeddingModel(FlagModel):
|
|
||||||
def __init__(self, config, **kwargs):
|
|
||||||
info = EMBED_MODEL_INFO[config.embed_model]
|
|
||||||
model_name_or_path = config.model_local_paths.get(info["name"], info.get("default_path"))
|
|
||||||
logger.info(f"Loading embedding model {info['name']} from {model_name_or_path}")
|
|
||||||
|
|
||||||
super().__init__(model_name_or_path,
|
|
||||||
query_instruction_for_retrieval=info.get("query_instruction", None),
|
|
||||||
use_fp16=False, **kwargs)
|
|
||||||
|
|
||||||
logger.info(f"Embedding model {info['name']} loaded")
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
from zhipuai import ZhipuAI
|
|
||||||
|
|
||||||
|
|
||||||
class RemoteEmbeddingModel:
|
class RemoteEmbeddingModel:
|
||||||
embed_state = {}
|
embed_state = {}
|
||||||
|
|
||||||
@ -51,6 +35,44 @@ class RemoteEmbeddingModel:
|
|||||||
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
class LocalEmbeddingModel(FlagModel, RemoteEmbeddingModel):
|
||||||
|
def __init__(self, config, **kwargs):
|
||||||
|
info = EMBED_MODEL_INFO[config.embed_model]
|
||||||
|
model_name_or_path = config.model_local_paths.get(info["name"], info.get("default_path"))
|
||||||
|
logger.info(f"Loading embedding model {info['name']} from {model_name_or_path}")
|
||||||
|
|
||||||
|
super().__init__(model_name_or_path,
|
||||||
|
query_instruction_for_retrieval=info.get("query_instruction", None),
|
||||||
|
use_fp16=False, **kwargs)
|
||||||
|
|
||||||
|
logger.info(f"Embedding model {info['name']} loaded")
|
||||||
|
|
||||||
|
|
||||||
|
def batch_encode(self, messages, batch_size=20):
|
||||||
|
logger.info(f"Batch encoding {len(messages)} messages")
|
||||||
|
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")
|
||||||
|
response = self.encode_queries(group_msg)
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
class ZhipuEmbedding(RemoteEmbeddingModel):
|
class ZhipuEmbedding(RemoteEmbeddingModel):
|
||||||
|
|
||||||
def __init__(self, config) -> None:
|
def __init__(self, config) -> None:
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user