From e998102eeb4815bae7fd4b7c14672bb0642544d7 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Mon, 3 Mar 2025 21:57:27 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E6=9C=AC=E5=9C=B0embedding?= =?UTF-8?q?=E6=A8=A1=E5=9E=8B=E7=9A=84=E8=BF=90=E8=A1=8Cbug=EF=BC=8C?= =?UTF-8?q?=E6=B7=BB=E5=8A=A0=E4=BA=86=20batch=5Fencode=20=E7=9A=84?= =?UTF-8?q?=E6=96=B9=E6=B3=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/models/embedding.py | 56 ++++++++++++++++++++++++++++------------- 1 file changed, 39 insertions(+), 17 deletions(-) diff --git a/src/models/embedding.py b/src/models/embedding.py index 853a9020..e7d7b6b4 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -2,28 +2,12 @@ import os import json import requests from FlagEmbedding import FlagModel +from zhipuai import ZhipuAI from src.config import EMBED_MODEL_INFO 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: embed_state = {} @@ -51,6 +35,44 @@ class RemoteEmbeddingModel: 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): def __init__(self, config) -> None: