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 @@
+
+
+
+
{{ message.model_name }}
+
+
+
+
+
+
+
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 @@