From 1bd09c7356e0db8154f5ad7e380937574376cf69 Mon Sep 17 00:00:00 2001 From: rootkit Date: Fri, 9 Jan 2026 04:51:27 +0800 Subject: [PATCH] update list api and embedding api --- PLUGINS/Embeddings/embeddings_qdrant.py | 16 ++++++++++++---- PLUGINS/Huggingface/download_model.py | 18 ++++++++---------- 2 files changed, 20 insertions(+), 14 deletions(-) diff --git a/PLUGINS/Embeddings/embeddings_qdrant.py b/PLUGINS/Embeddings/embeddings_qdrant.py index 12d4342..dd7929a 100644 --- a/PLUGINS/Embeddings/embeddings_qdrant.py +++ b/PLUGINS/Embeddings/embeddings_qdrant.py @@ -5,6 +5,7 @@ import os import uuid import httpx +import torch import urllib3 from langchain_classic.retrievers import ContextualCompressionRetriever from langchain_classic.retrievers.document_compressors import CrossEncoderReranker @@ -16,9 +17,13 @@ from langchain_qdrant import QdrantVectorStore, FastEmbedSparse, RetrievalMode from qdrant_client import models from Lib.configs import BASE_DIR +from Lib.log import logger from PLUGINS.Embeddings.CONFIG import EMBEDDINGS_TYPE, EMBEDDINGS_BASE_URL, EMBEDDINGS_MODEL, EMBEDDINGS_API_KEY, EMBEDDINGS_SIZE, EMBEDDINGS_PROXY from PLUGINS.Qdrant.qdrant import Qdrant +# 检查 GPU 是否可用 +device = "cuda" if torch.cuda.is_available() else "cpu" +logger.info(f"Using device for embeddings and reranker: {device}") urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning) @@ -32,7 +37,10 @@ class EmbeddingsAPI(object): self.dense_model = self.get_dense_model() self.vector_client = Qdrant.get_client() self.rerank_model = HuggingFaceCrossEncoder(model_name=os.path.join(BASE_DIR, 'Docker', 'Huggingface', 'bge-reranker-v2-m3'), - model_kwargs={'local_files_only': True}) + model_kwargs={ + 'local_files_only': True, + 'device': device + }) @staticmethod def get_dense_model(): @@ -104,13 +112,13 @@ class EmbeddingsAPI(object): ) return vector_store - def add_document(self, collection_name: str, ids: str, page_content: str, metadata: dict) -> list[str]: + def add_document(self, collection_name: str, ids: str, page_content: str, metadata: dict) -> str: namespace = uuid.NAMESPACE_DNS doc_id = str(uuid.uuid5(namespace, ids)) vector_store = self.vector_store(collection_name) document = Document(id=doc_id, page_content=page_content, metadata=metadata) result = vector_store.add_documents([document]) - return result + return result[0] def delete_document(self, collection_name: str, ids: str) -> bool | None: namespace = uuid.NAMESPACE_DNS @@ -124,7 +132,7 @@ class EmbeddingsAPI(object): results = vector_store.similarity_search_with_score(query, k=k) return results - def search_documents_with_rerank(self, collection_name: str, query: str, k: int = 20, top_n: int = 5) -> list[tuple[Document, float]]: + def search_documents_with_rerank(self, collection_name: str, query: str, k: int = 20, top_n: int = 5) -> list[Document]: """ 手动实现重排序以获取分数 """ diff --git a/PLUGINS/Huggingface/download_model.py b/PLUGINS/Huggingface/download_model.py index 14d508f..5090b3d 100644 --- a/PLUGINS/Huggingface/download_model.py +++ b/PLUGINS/Huggingface/download_model.py @@ -1,25 +1,23 @@ # uncomment the code to set a custom Hugging Face endpoint if needed. this must run on the top of the script import os -os.environ["HF_ENDPOINT"] = "https://hf-mirror.com" -# os.environ['HTTP_PROXY'] = 'http://127.0.0.1:7890' -# os.environ['HTTPS_PROXY'] = 'http://127.0.0.1:7890' - -import os +# os.environ["HF_ENDPOINT"] = "https://hf-mirror.com" +os.environ['HTTP_PROXY'] = 'http://127.0.0.1:7890' +os.environ['HTTPS_PROXY'] = 'http://127.0.0.1:7890' from fastembed import SparseTextEmbedding from huggingface_hub import snapshot_download - -from Lib.configs import BASE_DIR +from pathlib import Path if __name__ == "__main__": # download reranker model (BAAI/bge-reranker-v2-m3) model_id = "BAAI/bge-reranker-v2-m3" - local_dir = os.path.join(BASE_DIR, 'Docker', 'Huggingface', '../../Docker/Huggingface/bge-reranker-v2-m3') + script_path = Path(__file__).resolve() + project_root = script_path.parents[2] + local_dir = project_root / "Docker" / "Huggingface" / "bge-reranker-v2-m3" snapshot_download(repo_id=model_id, local_dir=local_dir) print("Reranker model downloaded to:", local_dir) # download sparse model (bm25) - cache_dir = os.path.join(BASE_DIR, 'Docker', 'Huggingface', '../../Docker/Huggingface/bm25') - + cache_dir = project_root / "Docker" / "Huggingface" / "bm25" model = SparseTextEmbedding(model_name="Qdrant/bm25", cache_dir=cache_dir) print("Sparse model downloaded to:", cache_dir)