From 0b2dd0179b58af3f8de49eecc23d313b0e63f0fb Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sun, 25 Aug 2024 20:29:24 +0800 Subject: [PATCH] sync --- src/config/__init__.py | 38 ++++++- src/core/database.py | 16 ++- src/core/knowledgebase.py | 9 +- src/core/retriever.py | 34 +++--- src/models/chat_model.py | 2 - src/models/embedding.py | 47 ++++----- src/utils/logging_config.py | 17 +-- src/views/common_view.py | 14 ++- src/views/database_view.py | 3 +- web/src/assets/base.css | 1 + web/src/components/ChatComponent.vue | 95 +---------------- web/src/components/DebugComponent.vue | 85 +++++++++++++++ web/src/components/RefsComponent.vue | 143 ++++++++++++++++++++++++++ web/src/layouts/AppLayout.vue | 30 +++++- web/src/router/index.js | 13 +++ web/src/views/DataBaseView.vue | 19 +++- web/src/views/SettingView.vue | 8 +- 17 files changed, 405 insertions(+), 169 deletions(-) create mode 100644 web/src/components/DebugComponent.vue create mode 100644 web/src/components/RefsComponent.vue diff --git a/src/config/__init__.py b/src/config/__init__.py index 4e28a553..b98038c4 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -51,7 +51,7 @@ class Config(SimpleConfig): ## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths 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("embed_model", default="zhipu-embedding-3", des="Embedding 模型", choices=list(EMBED_MODEL_INFO.keys())) self.add_item("reranker", default="bge-reranker-v2-m3", des="Re-Ranker 模型", choices=["bge-reranker-v2-m3"]) self.add_item("model_local_paths", default={}, des="本地模型路径") ### <<< 默认配置结束 @@ -92,22 +92,29 @@ class Config(SimpleConfig): """根据传入的文件覆盖掉默认配置""" 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: content = f.read() if content: - self.update(json.loads(content)) + local_config = json.loads(content) + local_config.pop("_config_items") + self.update(local_config) else: print(f"{self.filename} is empty.") + elif self.filename.endswith(".yaml"): with open(self.filename, 'r') as f: content = f.read() if content: - self.update(yaml.safe_load(content)) + local_config = yaml.safe_load(content) + local_config.pop("_config_items") + self.update(local_config) else: print(f"{self.filename} is empty.") else: logger.warning(f"Unknown config file type {self.filename}") + else: logger.warning(f"\n\n{'='*70}\n{'Config file not found':^70}\n{'You can custum your config in `' + self.filename + '`':^70}\n{'='*70}\n\n") @@ -170,7 +177,30 @@ MODEL_NAMES = { "llama3.1-8b-instruct", "llama3-8b-instruct", "llama3.1-405b-instruct", - "baichuan2-7b-chat-v1", "qwen2-0.5b-instruct" ] +} + + +EMBED_MODEL_INFO = { + "bge-large-zh-v1.5": SimpleConfig({ + "name": "bge-large-zh-v1.5", + "default_path": "BAAI/bge-large-zh-v1.5", + "dimension": 1024, + "query_instruction": "为这个句子生成表示以用于检索相关文章:", + }), + "zhipu-embedding-2": SimpleConfig({ + "name": "zhipu-embedding-2", + "default_path": "embedding-2", + "dimension": 1024, + }), + "zhipu-embedding-3": SimpleConfig({ + "name": "zhipu-embedding-3", + "default_path": "embedding-3", + "dimension": 2048, + }), +} + +RERANKER_LIST = { + "bge-reranker-v2-m3": "BAAI/bge-reranker-v2-m3", } \ No newline at end of file diff --git a/src/core/database.py b/src/core/database.py index ccdc834a..05b7b432 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -9,10 +9,11 @@ logger = setup_logger("DataBaseManager") class DataBaseLite: - def __init__(self, name, description, db_type, **kwargs) -> None: + def __init__(self, name, description, db_type, dimension=None, **kwargs) -> None: self.name = name self.description = description self.db_type = db_type + self.dimension = dimension self.db_id = kwargs.get("db_id", hashstr(name)) self.metaname = kwargs.get("metaname", f"{db_type[:1]}{hashstr(name)}") self.metadata = kwargs.get("metaname", {}) @@ -37,7 +38,8 @@ class DataBaseLite: "embed_model": self.embed_model, "metaname": self.metaname, "metadata": self.metadata, - "files": self.files + "files": self.files, + "dimension": self.dimension } def to_json(self): @@ -115,10 +117,14 @@ class DataBaseManager: else: return {"message": "Graph base not enabled", "graph": {}} - def create_database(self, database_name, description, db_type): - new_database = DataBaseLite(database_name, description, db_type, embed_model=self.config.embed_model) + def create_database(self, database_name, description, db_type, dimension): + new_database = DataBaseLite(database_name, + description, + db_type, + embed_model=self.config.embed_model, + dimension=dimension) - self.knowledge_base.add_collection(new_database.metaname) + self.knowledge_base.add_collection(new_database.metaname, dimension) self.data["databases"].append(new_database) self._save_databases() return self.get_databases() diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 7c0d5113..327384f0 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -18,7 +18,6 @@ class KnowledgeBase: self.client = MilvusClient(self.milvus_path) def _init_config(self, config): - self.vector_dim = 1024 # 暂时不知道这个和 embedding model 的 embedding 大小有什么关系 self.milvus_path = os.path.join(config.save_dir, "data/vector_base/milvus.db") os.makedirs(os.path.dirname(self.milvus_path), exist_ok=True) @@ -40,14 +39,14 @@ class KnowledgeBase: # collection["id"] = hashstr(collection_name) return collection - def add_collection(self, collection_name): + def add_collection(self, collection_name, dimension=None): if self.client.has_collection(collection_name=collection_name): logger.warning(f"Collection {collection_name} already exists, drop it") self.client.drop_collection(collection_name=collection_name) self.client.create_collection( collection_name=collection_name, - dimension=self.vector_dim, # The vectors we will use in this demo has 768 dimensions + dimension= dimension, # The vectors we will use in this demo has 768 dimensions ) def add_documents(self, docs, collection_name, **kwargs): @@ -55,8 +54,8 @@ class KnowledgeBase: # 检查 collection 是否存在 import random if not self.client.has_collection(collection_name=collection_name): - logger.warning(f"Collection {collection_name} not found, create it") - self.add_collection(collection_name) + logger.error(f"Collection {collection_name} not found, create it") + # self.add_collection(collection_name) vectors = self.embed_model.encode(docs) diff --git a/src/core/retriever.py b/src/core/retriever.py index 6f74c7f4..2025f3a2 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -66,29 +66,31 @@ class Retriever: def query_knowledgebase(self, query, history, refs): """查询知识库""" - rw_query = self.rewrite_query(query, history, refs) kb_res = [] final_res = [] - if refs["meta"].get("db_name") and self.config.enable_knowledge_base: + if not refs["meta"].get("db_name") or not self.config.enable_knowledge_base: + return {"results": final_res, "all_results": kb_res, "rw_query": query, "message": "Knowledge base is disabled"} - db_name = refs["meta"]["db_name"] - kb = self.dbm.metaname2db[db_name] - limit = refs["meta"].get("queryCount", 10) + rw_query = self.rewrite_query(query, history, refs) - kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=limit) + db_name = refs["meta"]["db_name"] + kb = self.dbm.metaname2db[db_name] + limit = refs["meta"].get("queryCount", 10) + + kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=limit) + for r in kb_res: + r["file"] = kb.id2file(r["entity"]["file_id"]) + + if self.config.enable_reranker: + RERANK_THRESHOLD = 0.1 for r in kb_res: - r["file"] = kb.id2file(r["entity"]["file_id"]) + r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True) + kb_res.sort(key=lambda x: x["rerank_score"], reverse=True) + final_res = [_res for _res in kb_res if _res["rerank_score"] > RERANK_THRESHOLD] - if self.config.enable_reranker: - RERANK_THRESHOLD = 0.1 - for r in kb_res: - r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True) - kb_res.sort(key=lambda x: x["rerank_score"], reverse=True) - final_res = [_res for _res in kb_res if _res["rerank_score"] > RERANK_THRESHOLD] - - else: - final_res = kb_res[:5] + else: + final_res = kb_res[:5] return {"results": final_res, "all_results": kb_res, "rw_query": rw_query} diff --git a/src/models/chat_model.py b/src/models/chat_model.py index da3c7765..be612215 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -11,8 +11,6 @@ class OpenAIBase(): self.model_name = model_name def predict(self, message, stream=False): - - logger.debug(message) if isinstance(message, str): messages=[{"role": "user", "content": message}] else: diff --git a/src/models/embedding.py b/src/models/embedding.py index c6246724..e7c482d9 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -1,37 +1,24 @@ import os from FlagEmbedding import FlagModel, FlagReranker +from src.config import EMBED_MODEL_INFO, RERANKER_LIST from src.utils.logging_config import setup_logger logger = setup_logger("EmbeddingModel") -SUPPORT_LIST = { - "bge-large-zh-v1.5": "BAAI/bge-large-zh-v1.5", - "zhipu": "embedding-3", -} - -RERANKER_LIST = { - "bge-reranker-v2-m3": "BAAI/bge-reranker-v2-m3", -} - -QUERY_INSTRUCTION = { - "bge-large-zh-v1.5": "为这个句子生成表示以用于检索相关文章:", -} class EmbeddingModel(FlagModel): - def __init__(self, config, **kwargs): - - assert config.embed_model in SUPPORT_LIST.keys(), f"Unsupported embed model: {config.embed_model}, only support {SUPPORT_LIST}" - - model_name_or_path = config.model_local_paths.get(config.embed_model, SUPPORT_LIST[config.embed_model]) - logger.info(f"Loading embedding model {config.embed_model} from {model_name_or_path}") + def __init__(self, model_info, config, **kwargs): + self.info = model_info + model_name_or_path = config.model_local_paths.get(model_info.name, model_info.default_path) + logger.info(f"Loading embedding model {model_info.name} from {model_name_or_path}") super().__init__(model_name_or_path, - query_instruction_for_retrieval=QUERY_INSTRUCTION[config.embed_model], + query_instruction_for_retrieval=model_info.get("query_instruction", None), use_fp16=False, **kwargs) - logger.info(f"Embedding model {config.embed_model} loaded") + logger.info(f"Embedding model {model_info.name} loaded") class Reranker(FlagReranker): @@ -50,8 +37,9 @@ from zhipuai import ZhipuAI class ZhipuEmbedding: - def __init__(self, config) -> None: + def __init__(self, model_info, config) -> None: self.config = config + self.model_info = model_info self.client = ZhipuAI(api_key=os.getenv("ZHIPUAPI")) logger.info("Zhipu Embedding model loaded") self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:" @@ -63,8 +51,8 @@ class ZhipuEmbedding: for i in range(0, len(message), 10): group_msg = message[i:i+10] response = self.client.embeddings.create( - model=SUPPORT_LIST[self.config.embed_model], - input=group_msg + model=self.model_info.default_path, + input=group_msg, ) data.extend([a.embedding for a in response.data]) @@ -83,7 +71,12 @@ def get_embedding_model(config): if not config.enable_knowledge_base: return None - if config.embed_model == "zhipu": - return ZhipuEmbedding(config) - else: - return EmbeddingModel(config) \ No newline at end of file + assert config.embed_model in EMBED_MODEL_INFO.keys(), f"Unsupported embed model: {config.embed_model}, only support {EMBED_MODEL_INFO.keys()}" + + if config.embed_model in ["bge-large-zh-v1.5"]: + model = EmbeddingModel(EMBED_MODEL_INFO[config.embed_model], config) + + if config.embed_model in ["zhipu-embedding-2", "zhipu-embedding-3"]: + model = ZhipuEmbedding(EMBED_MODEL_INFO[config.embed_model], config) + + return model \ No newline at end of file diff --git a/src/utils/logging_config.py b/src/utils/logging_config.py index 6f4849ce..2eb64b50 100644 --- a/src/utils/logging_config.py +++ b/src/utils/logging_config.py @@ -3,21 +3,23 @@ import os from datetime import datetime -# DATETIME = datetime.now().strftime('%Y-%m-%d-%H%M%S') -DATETIME = "debug" # 为了方便,调试的时候输出到 debug.log 文件 +DATETIME = datetime.now().strftime('%Y-%m-%d-%H%M%S') +# DATETIME = "debug" # 为了方便,调试的时候输出到 debug.log 文件 +LOG_FILE = f'log/project-{DATETIME}.log' -def setup_logger(name, log_file=None, level=logging.DEBUG, console=False): - - if log_file is None: - log_file = f'log/project-{DATETIME}.log' +def setup_logger(name, level=logging.DEBUG, console=False): os.makedirs("log", exist_ok=True) """Function to setup logger with the given name and log file.""" logger = logging.getLogger(name) logger.setLevel(level) + # 清除已有的 Handler,防止重复添加 + if logger.hasHandlers(): + logger.handlers.clear() + # File handler for logging to a file - file_handler = logging.FileHandler(log_file) + file_handler = logging.FileHandler(LOG_FILE) file_handler.setLevel(level) # Formatter for the logs @@ -34,6 +36,7 @@ def setup_logger(name, log_file=None, level=logging.DEBUG, console=False): return logger + # Setup the root logger logger = setup_logger('Athena') diff --git a/src/views/common_view.py b/src/views/common_view.py index d9445401..3eec200f 100644 --- a/src/views/common_view.py +++ b/src/views/common_view.py @@ -29,7 +29,6 @@ def chat(): request_data = json.loads(request.data) query = request_data['query'] meta = request_data.get('meta') - logger.debug(f"Web query: {query}") history_manager = HistoryManager(request_data['history']) new_query, refs = startup.retriever(query, history_manager.messages, meta) @@ -63,7 +62,7 @@ def call(): request_data = json.loads(request.data) query = request_data['query'] response = startup.model.predict(query) - logger.debug(f"\n\n\nCall query: \n{query} \n\nResponse: \n{response.content}\n\n") + logger.debug({"query": query, "response": response.content}) return jsonify({ "response": response.content, @@ -83,4 +82,13 @@ def update_config(): @common.route('/restart', methods=['POST']) def restart(): startup.restart() - return jsonify({"message": "Restarted!"}) \ No newline at end of file + return jsonify({"message": "Restarted!"}) + + +@common.route('/log', methods=['GET']) +def get_log(): + from src.utils.logging_config import LOG_FILE + with open(LOG_FILE, 'r') as f: + log = f.read() + + return jsonify({"log": log}) \ No newline at end of file diff --git a/src/views/database_view.py b/src/views/database_view.py index 29ec83f2..43bf1b75 100644 --- a/src/views/database_view.py +++ b/src/views/database_view.py @@ -27,8 +27,9 @@ def create_database(): database_name = data.get('database_name') description = data.get('description') db_type = data.get('db_type') + dimension = data.get('dimension') logger.debug(f"Create database {database_name}") - database = startup.dbm.create_database(database_name, description, db_type) + database = startup.dbm.create_database(database_name, description, db_type, dimension=dimension) return jsonify(database) @db.route('/', methods=['DELETE']) diff --git a/web/src/assets/base.css b/web/src/assets/base.css index 9e0174c1..dbc51192 100644 --- a/web/src/assets/base.css +++ b/web/src/assets/base.css @@ -10,6 +10,7 @@ --main-200: #8CC6E1; --main-100: #ABE0F7; --main-50: #CDF5FF; + --main-25: #E6FAFF; --c-white: #ffffff; --c-white-soft: #f8f8f8; diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index f2cf4348..93f51c8a 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -117,49 +117,7 @@ class="message-md" @click="consoleMsg(message)">

-
- - {{ filename }} - -
-

文件名: {{ results[0].file.filename }}

-

文件类型: {{ results[0].file.type }}

-

创建时间: {{ new Date(results[0].file.created_at * 1000).toLocaleString() }}

-
-
-

ID: #{{ res.id }}

-

- 相似度距离: -

- -
-

-

- 重排序分数: -

- -
-

- -

{{ res.entity.text }}

-
-
-
-
+
@@ -196,11 +154,14 @@ import { PlusCircleOutlined, FolderOutlined, FolderOpenOutlined, + GlobalOutlined, + FileTextOutlined, } from '@ant-design/icons-vue' import { onClickOutside } from '@vueuse/core' import { Marked } from 'marked'; import { markedHighlight } from 'marked-highlight'; import { useConfigStore } from '@/stores/config' +import RefsComponent from '@/components/RefsComponent.vue' import hljs from 'highlight.js'; import 'highlight.js/styles/github.css'; @@ -663,57 +624,9 @@ watch( margin-bottom: 0; } - .refs { - margin-bottom: 20px; - .filetag:hover { - cursor: pointer; - } - } } -.retrieval-detail { - .fileinfo { - margin-bottom: 20px; - padding: 1rem; - background: var(--main-50); - color: var(--main-800); - border-radius: 8px; - // border: 1px solid var(--main-100); - p { - margin: 10px; - line-height: 1.5; - } - } - - .result-item { - margin-bottom: 20px; - padding: 24px 16px 10px 16px; - border: 1px solid #e8e8e8; - border-radius: 8px; - background: var(--main-light-6); - - .result-id, - .result-distance, - .result-rerank-score, - .result-text-label, - .result-text { - margin: 5px 0; - } - - .scorebar { - margin-left: 10px; - display: inline-block; - width: 200px; - padding-bottom: 2px; - - - & > * { - margin: 0; - } - } - } -} diff --git a/web/src/components/DebugComponent.vue b/web/src/components/DebugComponent.vue new file mode 100644 index 00000000..19ac9500 --- /dev/null +++ b/web/src/components/DebugComponent.vue @@ -0,0 +1,85 @@ + + + + + diff --git a/web/src/components/RefsComponent.vue b/web/src/components/RefsComponent.vue new file mode 100644 index 00000000..3e83fae8 --- /dev/null +++ b/web/src/components/RefsComponent.vue @@ -0,0 +1,143 @@ + + + + + + diff --git a/web/src/layouts/AppLayout.vue b/web/src/layouts/AppLayout.vue index abaebc4f..8e6dc9f7 100644 --- a/web/src/layouts/AppLayout.vue +++ b/web/src/layouts/AppLayout.vue @@ -1,5 +1,5 @@