From 3ab20611a29b0fa6a3e8c937890c30db9476adad Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Mon, 22 Jul 2024 00:00:54 +0800 Subject: [PATCH] modified config embedding model and startup --- src/config/__init__.py | 25 ++++++++++--- src/core/database.py | 17 ++++++--- src/core/retriever.py | 13 +++---- src/core/startup.py | 23 +++++++++--- src/models/__init__.py | 2 +- src/models/embedding.py | 30 ++++++++++------ src/views/common_view.py | 25 ++++++++----- src/views/database_view.py | 21 ++++++----- web/src/layouts/AppLayout.vue | 18 ++++++++-- web/src/router/index.js | 2 +- web/src/stores/counter.js | 14 ++++++++ web/src/views/DataBaseInfoView.vue | 13 +++++-- web/src/views/DataBaseView.vue | 13 ++++--- web/src/views/SettingView.vue | 56 ++++++++++++++++++++++++++++++ 14 files changed, 212 insertions(+), 60 deletions(-) create mode 100644 web/src/views/SettingView.vue diff --git a/src/config/__init__.py b/src/config/__init__.py index a2d6f23f..2948cb7c 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -31,8 +31,7 @@ class Config(SimpleConfig): def __init__(self, filename=None): super().__init__() - self.filename = filename - logger.info(f"Loading config from {filename}") + self.filename = filename or "config/base.yaml" ### >>> 默认配置 # 可以在 config/base.yaml 中覆盖 @@ -48,6 +47,8 @@ class Config(SimpleConfig): # 模型配置 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths + self.model_provider = "qianfan" + self.model_name = None # 默认使用 provider 的默认模型 self.embed_model = "bge-large-zh-v1.5" self.reranker = "bge-reranker-v2-m3" ### <<< 默认配置结束 @@ -66,6 +67,7 @@ class Config(SimpleConfig): def load(self): """根据传入的文件覆盖掉默认配置""" + logger.info(f"Loading config from {self.filename}") if self.filename is not None and os.path.exists(self.filename): if self.filename.endswith(".json"): with open(self.filename, 'r') as f: @@ -77,7 +79,20 @@ class Config(SimpleConfig): logger.warning(f"Config file {self.filename} not found") def save(self): - if self.filename is not None: - with open(self.filename, 'w') as f: + logger.info(f"Saving config to {self.filename}") + if self.filename is None: + logger.warning("Config file is not specified, save to default config/base.yaml") + self.filename = "config/base.yaml" + + if self.filename.endswith(".json"): + with open(self.filename, 'w+') as f: + json.dump(self, f, indent=4, ensure_ascii=False) + elif self.filename.endswith(".yaml"): + with open(self.filename, 'w+') as f: + yaml.dump(self, f, indent=2) + else: + logger.warning(f"Unknown config file type {self.filename}, save as json") + with open(self.filename, 'w+') as f: json.dump(self, f, indent=4) - logger.info(f"Config file {self.filename} saved") + + logger.info(f"Config file {self.filename} saved") \ No newline at end of file diff --git a/src/core/database.py b/src/core/database.py index baefefbf..2de48798 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -6,7 +6,7 @@ from plugins import pdf2txt from core.knowledgebase import KnowledgeBase from core.filereader import pdfreader, plainreader from core.graphbase import GraphDatabase -from models.embedding import EmbeddingModel +from models.embedding import get_embedding_model logger = setup_logger("DataBaseManager") @@ -17,9 +17,10 @@ class DataBaseLite: self.description = description self.db_type = db_type self.db_id = kwargs.get("db_id", hashstr(name)) - self.metaname = kwargs.get("metaname", f"{db_type}_{hashstr(name)}") + self.metaname = kwargs.get("metaname", f"{db_type[:1]}{hashstr(name)}") self.metadata = kwargs.get("metaname", {}) self.files = kwargs.get("files", []) + self.embed_model = kwargs.get("embed_model", None) def update(self, metadata): @@ -31,6 +32,7 @@ class DataBaseLite: "description": self.description, "db_type": self.db_type, "db_id": self.db_id, + "embed_model": self.embed_model, "metaname": self.metaname, "metadata": self.metadata, "files": self.files @@ -47,7 +49,7 @@ class DataBaseManager: def __init__(self, config=None) -> None: self.config = config self.database_path = "data/databases.json" - self.embed_model = EmbeddingModel(config) + self.embed_model = get_embedding_model(config) self.knowledge_base = KnowledgeBase(config, self.embed_model) self.data = {"databases": [], "graph": {}} @@ -96,7 +98,7 @@ class DataBaseManager: return {"graph": {}, "message": "Graph database is not enabled"} def create_database(self, database_name, description, db_type): - new_database = DataBaseLite(database_name, description, db_type) + new_database = DataBaseLite(database_name, description, db_type, embed_model=self.config.embed_model) self.knowledge_base.add_collection(new_database.metaname) self.data["databases"].append(new_database) @@ -105,6 +107,11 @@ class DataBaseManager: def add_files(self, db_id, files): db = self.get_kb_by_id(db_id) + + if db.embed_model != self.config.embed_model: + logger.error(f"Embed model not match, {db.embed_model} != {self.config.embed_model}") + return {"message": "Embed model not match", "status": "failed"} + new_files = [] for file in files: # filenames = [f["filename"] for f in db.files] @@ -138,6 +145,8 @@ class DataBaseManager: self._save_databases() + return {"message": "全部解析完成", "status": "success"} + def get_database_info(self, db_id): db = self.get_kb_by_id(db_id) if db is None: diff --git a/src/core/retriever.py b/src/core/retriever.py index 8853fafb..4cc183e7 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -1,4 +1,3 @@ -from core.startup import dbm, model from models.embedding import Reranker from utils.logging_config import setup_logger logger = setup_logger("server-common") @@ -6,9 +5,11 @@ logger = setup_logger("server-common") class Retriever: - def __init__(self, config): + def __init__(self, config, dbm, model): self.config = config self.reranker = Reranker(config) + self.dbm = dbm + self.model = model def retrieval(self, query, history, meta): @@ -56,7 +57,7 @@ class Retriever: results = [] if meta.get("use_graph"): for entitie in entities: - result = dbm.graph_base.query_by_vector(entitie) + result = self.dbm.graph_base.query_by_vector(entitie) if result != []: results.extend(result) return {"results": self.format_query_results(results)} @@ -65,7 +66,7 @@ class Retriever: kb_res = [] if meta.get("db_name"): - kb_res = dbm.knowledge_base.search(query, meta["db_name"], limit=5) + kb_res = self.dbm.knowledge_base.search(query, meta["db_name"], limit=5) for r in kb_res: r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True) @@ -97,7 +98,7 @@ class Retriever: # 构建提示词 rewritten_query_prompt = rewritten_query_prompt_template.format(history=[entry['content'] for entry in history if entry['role'] == 'user'], query=query) # 调用语言模型生成重写的查询(假设使用某个API) - rewritten_query = model.predict(rewritten_query_prompt).content + rewritten_query = self.model.predict(rewritten_query_prompt).content if meta.get("use_graph"): @@ -113,7 +114,7 @@ class Retriever: """ # 构建提示词 entity_extraction_prompt = entity_extraction_prompt_template.format(text=rewritten_query) - entities = model.predict(entity_extraction_prompt).content.split(",") + entities = self.model.predict(entity_extraction_prompt).content.split(",") entities = [entity for entity in entities if all(char.isalnum() or char in '汉字' for char in entity)] else: entities = [] diff --git a/src/core/startup.py b/src/core/startup.py index 70d0ae7a..4a862028 100644 --- a/src/core/startup.py +++ b/src/core/startup.py @@ -1,10 +1,25 @@ from core import DataBaseManager +from core.retriever import Retriever from models import select_model from config import Config +from utils import setup_logger + +logger = setup_logger("Startup") -config = Config("config/base.yaml") -model = select_model(config) -dbm = DataBaseManager(config) +class Startup: + def __init__(self): + self.config = Config("config/base.yaml") + self.model = select_model(self.config) + self.dbm = DataBaseManager(self.config) + self.retriever = Retriever(self.config, self.dbm, self.model) -# 启动本地图数据库 \ No newline at end of file + def restart(self): + logger.info("Restarting...") + self.model = select_model(self.config) + self.dbm = DataBaseManager(self.config) + self.retriever = Retriever(self.config, self.dbm, self.model) + logger.info("Restarted") + + +startup = Startup() \ No newline at end of file diff --git a/src/models/__init__.py b/src/models/__init__.py index 35cd0c5d..33142937 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -6,7 +6,7 @@ def select_model(config): model_provider = config.model_provider model_name = config.model_name - logger.info(f"Selecting model from {model_provider} with name {model_name}") + logger.info(f"Selecting model from {model_provider} with {model_name or 'default'}") if model_provider == "deepseek": from models.chat_model import DeepSeek diff --git a/src/models/embedding.py b/src/models/embedding.py index 6dc11006..3097dfc8 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -40,31 +40,39 @@ class Reranker(FlagReranker): assert config.reranker in RERANKER_LIST.keys(), f"Unsupported Reranker: {config.reranker}, only support {RERANKER_LIST.keys()}" model_name_or_path = config.model_local_paths.get(config.reranker, RERANKER_LIST[config.reranker]) - logger.info(f"Loading Reranker model {config.re_ranker} from {model_name_or_path}") + logger.info(f"Loading Reranker model {config.reranker} from {model_name_or_path}") super().__init__(model_name_or_path, use_fp16=True, **kwargs) - logger.info(f"Reranker model {config.re_ranker} loaded") + logger.info(f"Reranker model {config.reranker} loaded") from zhipuai import ZhipuAI -client = ZhipuAI(api_key="270ea71e9560c0ff406acbcdd48bfd97.e3XOMdWKuZb7Q1Sk") -response = client.embeddings.create( - model="embedding-2", #填写需要调用的模型名称 - input=["你好","woshi"] -) - -print(response.data.shape) - class ZhipuEmbedding: def __init__(self, config) -> None: self.config = config self.client = ZhipuAI(api_key=os.getenv("ZHIPUAPI")) + logger.info("Zhipu Embedding model loaded") + self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:" def predict(self, message): response = self.client.embeddings.create( model=SUPPORT_LIST[self.config.embed_model], input=message ) - return response.data \ No newline at end of file + return [a["embedding"] for a in response["data"]] + + def encode(self, message): + return self.predict(message) + + def encode_queries(self, queries): + # queries = [self.query_instruction_for_retrieval + query for query in queries] + return self.predict(queries) + + +def get_embedding_model(config): + if config.embed_model == "zhipu": + return ZhipuEmbedding(config) + else: + return EmbeddingModel(config) \ No newline at end of file diff --git a/src/views/common_view.py b/src/views/common_view.py index 0436ead6..58d0d04f 100644 --- a/src/views/common_view.py +++ b/src/views/common_view.py @@ -3,13 +3,10 @@ from flask import Blueprint, jsonify, request, Response from core import HistoryManager from utils.logging_config import setup_logger -from core.startup import config, model -from core.retriever import Retriever - +from core.startup import startup common = Blueprint('common', __name__) logger = setup_logger("server-common") -retriever = Retriever(config) @common.route('/', methods=["GET"]) def route_index(): @@ -35,7 +32,7 @@ def chat(): logger.debug(f"Web query: {query}") history_manager = HistoryManager(request_data['history']) - new_query, refs = retriever(query, history_manager.messages, meta) + new_query, refs = startup.retriever(query, history_manager.messages, meta) messages = history_manager.get_history_with_msg(new_query) history_manager.add_user(query) @@ -43,7 +40,7 @@ def chat(): def generate_response(): content = "" - for delta in model.predict(messages, stream=True): + for delta in startup.model.predict(messages, stream=True): if delta.content: content += delta.content response_chunk = json.dumps({ @@ -59,7 +56,7 @@ def chat(): def call(): request_data = json.loads(request.data) query = request_data['query'] - response = model.predict(query) + response = startup.model.predict(query) logger.debug(f"Call query: {query} Response: {response.content}") return jsonify({ @@ -68,4 +65,16 @@ def call(): @common.route('/config', methods=['get']) def get_config(): - return jsonify(config) \ No newline at end of file + return jsonify(startup.config) + +@common.route('/config', methods=['post']) +def update_config(): + request_data = json.loads(request.data) + startup.config.update(request_data) + startup.config.save() + return jsonify(startup.config) + +@common.route('/restart', methods=['POST']) +def restart(): + startup.restart() + return jsonify({"message": "Restarted!"}) \ No newline at end of file diff --git a/src/views/database_view.py b/src/views/database_view.py index 30a26d89..a321c4b0 100644 --- a/src/views/database_view.py +++ b/src/views/database_view.py @@ -3,9 +3,8 @@ import json import threading from flask import Blueprint, jsonify, request, Response -from core import HistoryManager from utils.logging_config import setup_logger -from core.startup import config, model, dbm +from core.startup import startup db = Blueprint('database', __name__, url_prefix="/database") @@ -15,7 +14,7 @@ progress = {} # 只针对单个用户的进度 @db.route('/', methods=['GET']) def get_databases(): - database = dbm.get_databases() + database = startup.dbm.get_databases() return jsonify(database) @db.route('/', methods=['POST']) @@ -25,7 +24,7 @@ def create_database(): description = data.get('description') db_type = data.get('db_type') logger.debug(f"Create database {database_name}") - database = dbm.create_database(database_name, description, db_type) + database = startup.dbm.create_database(database_name, description, db_type) return jsonify(database) # TODO: 删除数据库 @@ -34,7 +33,7 @@ def delete_database(): data = json.loads(request.data) db_id = data.get('db_id') logger.debug(f"Delete database {db_id}") - dbm.delete_database(db_id) + startup.dbm.delete_database(db_id) return jsonify({"message": "删除成功"}) @@ -44,8 +43,8 @@ def create_document_by_file(): db_id = data.get('db_id') files = data.get('files') logger.debug(f"Add document in {db_id} by file: {files}") - dbm.add_files(db_id, files) - return jsonify({"status": "全部解析完成"}) + msg = startup.dbm.add_files(db_id, files) + return jsonify(msg) @db.route('/info', methods=['GET']) @@ -55,7 +54,7 @@ def get_database_info(): return jsonify({"message": "db_id is required"}), 400 logger.debug(f"Get database {db_id} info") - database = dbm.get_database_info(db_id) + database = startup.dbm.get_database_info(db_id) if database is None: return jsonify({"message": "database not found"}), 404 @@ -69,7 +68,7 @@ def delete_document(): db_id = data.get('db_id') file_id = data.get('file_id') logger.debug(f"DELETE document {file_id} info in {db_id}") - dbm.delete_file(db_id, file_id) + startup.dbm.delete_file(db_id, file_id) return jsonify({"message": "删除成功"}) @db.route('/document', methods=['GET']) @@ -77,7 +76,7 @@ def get_document_info(): db_id = request.args.get('db_id') file_id = request.args.get('file_id') logger.debug(f"GET document {file_id} info in {db_id}") - info = dbm.get_file_info(db_id, file_id) + info = startup.dbm.get_file_info(db_id, file_id) return jsonify(info) @db.route('/upload', methods=['POST']) @@ -97,5 +96,5 @@ def upload_file(): @db.route('/graph', methods=['GET']) def get_graph_info(): - graph_info = dbm.get_graph() + graph_info = startup.dbm.get_graph() return jsonify(graph_info) diff --git a/web/src/layouts/AppLayout.vue b/web/src/layouts/AppLayout.vue index 7f4365d4..bd926c1e 100644 --- a/web/src/layouts/AppLayout.vue +++ b/web/src/layouts/AppLayout.vue @@ -1,5 +1,5 @@ + +