From 0af178bf91b556aec5fd0956253944b3a16a8e21 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Tue, 11 Mar 2025 16:26:55 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9B=B4=E6=96=B0=E6=A8=A1=E5=9E=8B=E6=9C=AC?= =?UTF-8?q?=E5=9C=B0=E9=85=8D=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/models/embedding.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/src/models/embedding.py b/src/models/embedding.py index 58fc57f5..ea10bde9 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -33,7 +33,8 @@ class BaseEmbeddingModel: 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) + response = self.encode(group_msg) + logger.debug(f"Response: {len(response)=}, {len(group_msg)=}, {len(response[0])=}") data.extend(response) if len(messages) > batch_size: @@ -49,6 +50,9 @@ class LocalEmbeddingModel(FlagModel, BaseEmbeddingModel): self.model = config.model_local_paths.get(info["name"], info.get("local_path")) self.model = self.model or info["name"] + if os.path.exists(_path := os.path.join(os.getenv("MODEL_DIR"), self.model)): + self.model = _path + logger.info(f"Loading local model `{info['name']}` from `{self.model}` with device `{config.device}`") super().__init__(self.model,