From adf130740bbd2c7d4fd804b14498a23c5b914d83 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Wed, 7 Aug 2024 17:24:46 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9C=80=E5=B0=8F=E5=8C=96=E5=90=AF=E5=8A=A8?= =?UTF-8?q?=E4=BE=9D=E8=B5=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/config/__init__.py | 2 +- src/core/database.py | 6 +++--- src/core/filereader.py | 3 +-- src/core/knowledgebase.py | 1 + src/models/chat_model.py | 2 +- src/requirements.txt | 3 --- 6 files changed, 7 insertions(+), 10 deletions(-) diff --git a/src/config/__init__.py b/src/config/__init__.py index 24d450a5..dbd9f8a5 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -49,7 +49,7 @@ class Config(SimpleConfig): # 模型配置 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths - self.add_item("model_provider", default="qianfan", des="模型提供商", choices=["qianfan", "vllm", "zhipu", "deepseek", "dashscope"]) + self.add_item("model_provider", default="zhipu", des="模型提供商", choices=["qianfan", "vllm", "zhipu", "deepseek", "dashscope"]) self.add_item("model_name", default=None, des="模型名称") self.add_item("embed_model", default="bge-large-zh-v1.5", des="Embedding 模型", choices=["bge-large-zh-v1.5", "zhipu"]) self.add_item("reranker", default="bge-reranker-v2-m3", des="Re-Ranker 模型", choices=["bge-reranker-v2-m3"]) diff --git a/src/core/database.py b/src/core/database.py index 58751e40..c8dd43d8 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -2,10 +2,7 @@ import os import json import time from src.utils import hashstr, setup_logger, is_text_pdf -from src.plugins import pdf2txt -from src.core.knowledgebase import KnowledgeBase from src.core.filereader import pdfreader, plainreader -from src.core.graphbase import GraphDatabase from src.models.embedding import get_embedding_model logger = setup_logger("DataBaseManager") @@ -57,6 +54,8 @@ class DataBaseManager: self.embed_model = get_embedding_model(config) if self.config.enable_knowledge_base: + from src.core.knowledgebase import KnowledgeBase + from src.core.graphbase import GraphDatabase self.knowledge_base = KnowledgeBase(config, self.embed_model) self.graph_base = GraphDatabase(self.config, self.embed_model) self.graph_base.start() @@ -180,6 +179,7 @@ class DataBaseManager: if is_text_pdf(file): return pdfreader(file) else: + from src.plugins import pdf2txt return pdf2txt(file, return_text=True) elif file.endswith(".txt") or file.endswith(".md"): diff --git a/src/core/filereader.py b/src/core/filereader.py index ae9034dc..62fd11a6 100644 --- a/src/core/filereader.py +++ b/src/core/filereader.py @@ -1,14 +1,13 @@ import os from pathlib import Path -from llama_index.readers.file import PDFReader - def pdfreader(file_path): """读取PDF文件并返回text文本""" assert os.path.exists(file_path), "File not found" assert file_path.endswith(".pdf"), "File format not supported" + from llama_index.readers.file import PDFReader doc = PDFReader().load_data(file=Path(file_path)) # 简单的拼接起来之后返回纯文本 diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 20eb68f0..27ac52ab 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -14,6 +14,7 @@ class KnowledgeBase: assert embed_model, "embed_model=None" self.embed_model = embed_model + self.client = MilvusClient(self.milvus_path) def _init_config(self, config): diff --git a/src/models/chat_model.py b/src/models/chat_model.py index 286fe5d0..da3c7765 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -64,7 +64,6 @@ class VLLM(OpenAIBase): super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) -import qianfan class GeneralResponse: @@ -76,6 +75,7 @@ class GeneralResponse: class Qianfan: def __init__(self, model_name="ernie_speed") -> None: + import qianfan self.model_name = model_name access_key = os.getenv("QIANFAN_ACCESS_KEY") secret_key = os.getenv("QIANFAN_SECRET_KEY") diff --git a/src/requirements.txt b/src/requirements.txt index 32f414e5..29577253 100644 --- a/src/requirements.txt +++ b/src/requirements.txt @@ -1,9 +1,6 @@ FlagEmbedding==1.2.10 Flask==3.0.3 Flask_Cors==4.0.1 -llama_index==0.10.53 openai==1.35.10 -pymilvus==2.4.4 python-dotenv==1.0.1 PyYAML==6.0.1 -qianfan==0.4.0.1