diff --git a/src/config/__init__.py b/src/config/__init__.py
index 807f3125..24d450a5 100644
--- a/src/config/__init__.py
+++ b/src/config/__init__.py
@@ -43,15 +43,14 @@ class Config(SimpleConfig):
self.add_item("save_dir", default="saves", des="保存目录")
# 功能选项
self.add_item("enable_reranker", default=False, des="是否开启重排序")
- self.add_item("enable_knowledge_base", default=True, des="是否开启知识库")
- self.add_item("enable_knowledge_graph", default=True, des="是否开启知识图谱")
- self.add_item("enable_search_engine", default=True, des="是否开启搜索引擎")
+ self.add_item("enable_knowledge_base", default=False, des="是否开启知识库")
+ self.add_item("enable_search_engine", default=False, des="是否开启搜索引擎")
# 模型配置
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths
- self.add_item("model_provider", default="qianfan", des="模型提供商", choices=["qianfan", "vllm", "zhipu", "deepseek"])
- self.add_item("model_name", default=None, des="模型名称,为空则表示使用默认值")
+ self.add_item("model_provider", default="qianfan", 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"])
self.add_item("model_local_paths", default={}, des="本地模型路径")
@@ -80,6 +79,15 @@ class Config(SimpleConfig):
if not model_rel_path.startswith("/"):
self.model_local_paths[model] = os.path.join(model_root_dir, model_rel_path)
+ self.model_names = MODEL_NAMES
+
+ if self.model_name not in self.model_names[self.model_provider]:
+ logger.warning(f"Model name {self.model_name} not in {self.model_provider}, using default model name")
+ self.model_name = self.model_names[self.model_provider][0]
+
+ default_model_name = self.model_names[self.model_provider][0]
+ self.model_name = self.get("model_name") or default_model_name
+
def load(self):
"""根据传入的文件覆盖掉默认配置"""
logger.info(f"Loading config from {self.filename}")
@@ -121,4 +129,48 @@ class Config(SimpleConfig):
with open(self.filename, 'w+') as f:
json.dump(self, f, indent=4)
- logger.info(f"Config file {self.filename} saved")
\ No newline at end of file
+ logger.info(f"Config file {self.filename} saved")
+
+MODEL_NAMES = {
+ # https://platform.deepseek.com/api-docs/zh-cn/pricing
+ "deepseek": [
+ "deepseek-chat",
+ "deepseek-coder"
+ ],
+
+ # https://open.bigmodel.cn/dev/api glm-4-0520、glm-4 、glm-4-air、glm-4-airx、 glm-4-flash
+ "zhipu": [
+ "glm-4",
+ "glm-4-0520",
+ "glm-4-air",
+ "glm-4-airx",
+ "glm-4-flash"
+ ],
+
+ # {'ERNIE-4.0-8K-0104', 'ERNIE-Lite-8K-0308', 'ERNIE-Speed-128K', 'ERNIE-3.5-128K(预览版)', 'Yi-34B-Chat', 'ERNIE-4.0-8K-Preview-0518', 'ERNIE-Bot-4', 'ERNIE-3.5-128K', 'ChatGLM2-6B-32K', 'ERNIE-3.5-8K', 'EB-turbo-AppBuilder', 'ERNIE-Lite-AppBuilder-8K', 'ERNIE-4.0-8K-0329', 'AquilaChat-7B', 'Gemma-7B-it', 'Qianfan-Chinese-Llama-2-70B', 'Mixtral-8x7B-Instruct', 'Gemma-7B-It', 'ERNIE Speed-AppBuilder', 'ERNIE-Function-8K', 'ERNIE-4.0-8K-preview', 'ERNIE-Bot', 'Qianfan-BLOOMZ-7B-compressed', 'ERNIE-4.0-8K', 'BLOOMZ-7B', 'ERNIE-Character-8K', 'ERNIE-3.5-8K-0205', 'ERNIE-4.0-8K-0613', 'Llama-2-70B-Chat', 'ERNIE-Character-Fiction-8K', 'ERNIE-4.0-8K-Preview', 'ERNIE-3.5-8K-Preview', 'ERNIE-Speed', 'ERNIE-Tiny-8K', 'ERNIE-4.0-Turbo-8K-Preview', 'Meta-Llama-3-8B', 'ERNIE-4.0-8K-Latest', 'ERNIE 3.5', 'XuanYuan-70B-Chat-4bit', 'Llama-2-13B-Chat', 'ERNIE-Bot-turbo', 'ERNIE-3.5-8K-0613', 'ERNIE-Lite-AppBuilder-8K-0614', 'ERNIE-4.0-preview', 'Llama-2-7B-Chat', 'Qianfan-Chinese-Llama-2-13B', 'ERNIE-Bot-turbo-AI', 'Meta-Llama-3-70B', 'ERNIE-Functions-8K', 'ERNIE-Lite-8K-0922(原ERNIE-Bot-turbo-0922)', 'ERNIE Speed', 'ERNIE-3.5-preview', 'Qianfan-Chinese-Llama-2-7B', 'ERNIE-Speed-8K', 'ERNIE-Lite-8K-0922', 'ChatLaw', 'ERNIE-3.5-8K-0329', 'ERNIE-4.0-Turbo-8K', 'ERNIE-3.5-8K-preview', 'ERNIE-Lite-8K'}
+ "qianfan": [
+ "ERNIE-Speed",
+ "ERNIE-Speed-8K",
+ "ERNIE-Speed-128K",
+ "ERNIE-Tiny-8K",
+ "ERNIE-Lite-8K",
+ "ERNIE-4.0-8K-Latest"
+ "Yi-34B-Chat",
+ ],
+
+ "vllm": [
+ "vllm",
+ ],
+
+ # https://bailian.console.aliyun.com/?switchAgent=10226727&productCode=p_efm#/model-market
+ "dashscope": [
+ "qwen-long",
+ "qwen2-7b-instruct",
+ "qwen2-1.5b-instruct",
+ "llama3.1-8b-instruct",
+ "llama3-8b-instruct",
+ "llama3.1-405b-instruct",
+ "baichuan2-7b-chat-v1",
+ "qwen2-0.5b-instruct"
+ ]
+}
\ No newline at end of file
diff --git a/src/core/database.py b/src/core/database.py
index 6b951c52..58751e40 100644
--- a/src/core/database.py
+++ b/src/core/database.py
@@ -55,13 +55,14 @@ class DataBaseManager:
self.config = config
self.database_path = os.path.join(config.save_dir, "data", "database.json")
self.embed_model = get_embedding_model(config)
- self.knowledge_base = KnowledgeBase(config, self.embed_model)
- self.graph_base = GraphDatabase(self.config, self.embed_model)
- self.data = {"databases": [], "graph": {}}
- if self.config.enable_knowledge_graph:
+ if self.config.enable_knowledge_base:
+ self.knowledge_base = KnowledgeBase(config, self.embed_model)
+ self.graph_base = GraphDatabase(self.config, self.embed_model)
self.graph_base.start()
+ self.data = {"databases": [], "graph": {}}
+
self._load_databases()
self._update_database()
@@ -103,11 +104,11 @@ class DataBaseManager:
return {"databases": [db.to_dict() for db in self.data["databases"]]}
def get_graph(self):
- if self.config.enable_knowledge_graph:
+ if self.config.enable_graph_base:
self.data["graph"].update(self.graph_base.get_database_info("neo4j"))
return {"graph": self.data["graph"]}
else:
- return {"graph": {}, "message": "Graph database is not enabled"}
+ 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)
diff --git a/src/core/retriever.py b/src/core/retriever.py
index 50677467..19a93dff 100644
--- a/src/core/retriever.py
+++ b/src/core/retriever.py
@@ -59,7 +59,7 @@ class Retriever:
# res = model.predict("qiansdgsa, dasdh ashdsakjdk ak ").content
results = []
- if refs["meta"].get("use_graph"):
+ if refs["meta"].get("use_graph") and self.config.enable_knowledge_base:
for entity in refs["entities"]:
result = self.dbm.graph_base.query_by_vector(entity)
if result != []:
@@ -71,7 +71,8 @@ class Retriever:
query = refs.get("rewritten_query", query)
kb_res = []
- if refs["meta"].get("db_name"):
+ final_res = []
+ if refs["meta"].get("db_name") and self.config.enable_knowledge_base:
db_name = refs["meta"]["db_name"]
kb = self.dbm.metaname2db[refs["meta"]["db_name"]]
limit = refs["meta"].get("queryCount", 10)
diff --git a/src/core/startup.py b/src/core/startup.py
index bc85bba6..943dd10b 100644
--- a/src/core/startup.py
+++ b/src/core/startup.py
@@ -9,6 +9,9 @@ logger = setup_logger("Startup")
class Startup:
def __init__(self):
+ self.start()
+
+ def start(self):
self.config = Config()
self.model = select_model(self.config)
self.dbm = DataBaseManager(self.config)
@@ -16,9 +19,7 @@ class Startup:
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)
+ self.start()
logger.info("Restarted")
diff --git a/src/models/__init__.py b/src/models/__init__.py
index f2bb4c96..857e0b80 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 {model_name or 'default'}")
+ logger.info(f"Selecting model from {model_provider} with {model_name}")
if model_provider == "deepseek":
from src.models.chat_model import DeepSeek
@@ -24,6 +24,10 @@ def select_model(config):
from src.models.chat_model import VLLM
return VLLM(model_name)
+ elif model_provider == "dashscope":
+ from src.models.chat_model import DashScope
+ return DashScope(model_name)
+
elif model_provider is None:
raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`")
else:
diff --git a/src/models/chat_model.py b/src/models/chat_model.py
index a5b2a771..286fe5d0 100644
--- a/src/models/chat_model.py
+++ b/src/models/chat_model.py
@@ -67,9 +67,10 @@ class VLLM(OpenAIBase):
import qianfan
-class QianfanResponse:
+class GeneralResponse:
def __init__(self, content):
self.content = content
+ self.is_full = False
class Qianfan:
@@ -98,7 +99,7 @@ class Qianfan:
stream=True,
)
for chunk in response:
- yield QianfanResponse(chunk["body"]["result"])
+ yield GeneralResponse(chunk["body"]["result"])
def _get_response(self, messages):
response = self.client.do(
@@ -106,6 +107,55 @@ class Qianfan:
messages=messages,
stream=False,
)
- return QianfanResponse(response["body"]["result"])
+ return GeneralResponse(response["body"]["result"])
+
+class DashScope:
+
+ def __init__(self, model_name="qwen-long") -> None:
+ self.model_name = model_name
+ self.api_key= os.getenv("DASHSCOPE_API_KEY")
+
+
+ def predict(self, message, stream=False):
+ if isinstance(message, str):
+ messages=[{"role": "user", "content": message}]
+ else:
+ messages = message
+
+ if stream:
+ return self._stream_response(messages)
+ else:
+ return self._get_response(messages)
+
+ def _stream_response(self, messages):
+ import dashscope
+ response = dashscope.Generation.call(
+ api_key=self.api_key,
+ model=self.model_name,
+ messages=messages,
+ result_format='message',
+ stream=True,
+ )
+ for chunk in response:
+ message = chunk.output.choices[0].message
+ message.is_full = True
+ yield chunk.output.choices[0].message
+
+ def _get_response(self, messages):
+ import dashscope
+ response = dashscope.Generation.call(
+ api_key=self.api_key,
+ model=self.model_name,
+ messages=messages,
+ result_format='message',
+ stream=False,
+ )
+ return response.output.choices[0].message
+
+
+if __name__ == "__main__":
+ model = DashScope()
+ for a in model.predict("你好", stream=True):
+ print(a.content)
\ No newline at end of file
diff --git a/src/models/embedding.py b/src/models/embedding.py
index 06a4d8ea..64c6e173 100644
--- a/src/models/embedding.py
+++ b/src/models/embedding.py
@@ -72,6 +72,9 @@ class ZhipuEmbedding:
def get_embedding_model(config):
+ if not config.enable_knowledge_base:
+ return None
+
if config.embed_model == "zhipu":
return ZhipuEmbedding(config)
else:
diff --git a/src/views/common_view.py b/src/views/common_view.py
index 04bbc1e9..221264ad 100644
--- a/src/views/common_view.py
+++ b/src/views/common_view.py
@@ -41,14 +41,20 @@ def chat():
def generate_response():
content = ""
for delta in startup.model.predict(messages, stream=True):
- if delta.content:
+ if not delta.content:
+ continue
+
+ if hasattr(delta, 'is_full') and delta.is_full:
+ content = delta.content
+ else:
content += delta.content
- response_chunk = json.dumps({
- "history": history_manager.update_ai(content),
- "response": content,
- "refs": refs # TODO: 优化 refs,不需要每次都返回
- }, ensure_ascii=False).encode('utf8') + b'\n'
- yield response_chunk
+
+ response_chunk = json.dumps({
+ "history": history_manager.update_ai(content),
+ "response": content,
+ "refs": refs # TODO: 优化 refs,不需要每次都返回
+ }, ensure_ascii=False).encode('utf8') + b'\n'
+ yield response_chunk
return Response(generate_response(), content_type='application/json', status=200)
@@ -57,7 +63,7 @@ def call():
request_data = json.loads(request.data)
query = request_data['query']
response = startup.model.predict(query)
- logger.debug(f"Call query: {query} Response: {response.content}")
+ logger.debug(f"\n\n\nCall query: \n{query} \n\nResponse: \n{response.content}\n\n")
return jsonify({
"response": response.content,
diff --git a/src/views/database_view.py b/src/views/database_view.py
index d3069016..8516cdb9 100644
--- a/src/views/database_view.py
+++ b/src/views/database_view.py
@@ -14,7 +14,11 @@ progress = {} # 只针对单个用户的进度
@db.route('/', methods=['GET'])
def get_databases():
- database = startup.dbm.get_databases()
+ try:
+ database = startup.dbm.get_databases()
+ except Exception as e:
+ return jsonify({"message": "获取数据库列表失败", "databases": []})
+
return jsonify(database)
@db.route('/', methods=['POST'])
diff --git a/web/src/assets/base.css b/web/src/assets/base.css
index c6c658df..3d4e5f0e 100644
--- a/web/src/assets/base.css
+++ b/web/src/assets/base.css
@@ -18,7 +18,7 @@
--c-text-light-1: var(--c-indigo);
--c-text-light-2: rgba(60, 66, 70, 0.66);
--c-text-dark-1: var(--c-white);
- --c-text-dark-2: rgba(235, 235, 235, 0.64);
+ --c-text-dark-2: #b8b8b8;
}
/* semantic color variables for this project */
diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue
index 0ac32c5a..cdfae57d 100644
--- a/web/src/components/ChatComponent.vue
+++ b/web/src/components/ChatComponent.vue
@@ -43,7 +43,7 @@
即便强如雅典娜也可能会出错,请注意辨别内容的可靠性 模型供应商:{{ configStore.config?.model_provider }}
+即便强如雅典娜也可能会出错,请注意辨别内容的可靠性 模型供应商:{{ configStore.config?.model_provider }}:{{ configStore.config?.model_name }}
知识型数据库,主要是非结构化的文本组成,使用向量检索使用。