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 @@
+
+