diff --git a/.gitignore b/.gitignore
index b026948c..409b49b2 100644
--- a/.gitignore
+++ b/.gitignore
@@ -25,13 +25,19 @@ cache
### IDE
.vscode
+.idea
*.nogit.*
*.pdf
+*.yaml
src/data
neo4j*
*/package-lock.json
web/package-lock.json
saves
notebooks
-*.yaml
\ No newline at end of file
+local_neo4j/data
+local_neo4j/logs
+local_neo4j/import
+local_neo4j/plugins
+local_neo4j/conf
\ No newline at end of file
diff --git a/README.md b/README.md
index a391ec99..f4c8443e 100644
--- a/README.md
+++ b/README.md
@@ -3,22 +3,42 @@
-### 准备
+## 准备
1. 提供 API 服务商的 API_KEY,并放置在 `src/.env` 文件中,参考 `src/.env.template`。默认使用的是智谱AI。
-2. 配置 python 环境 `pip install -r src/requirements.txt`
+2. 配置 python 环境 `pip install -r requirements.txt`
+**如果不启用知识库,可以仅安装下面的依赖**
-### 启动命令行模式
-
-```bash
-python -m src.cli
+```
+FlagEmbedding==1.2.10
+Flask==3.0.3
+Flask_Cors==4.0.1
+openai==1.35.10
+python-dotenv==1.0.1
+PyYAML==6.0.1
+zhipuai
```
-### 启动网页模式
+### 【可选】配置图数据库 neo4j
+
+使用 docker 部署 neo4j 服务,配置文件见 [local_neo4j/docker-compose.yml](local_neo4j/docker-compose.yml).
+默认账号密码见最后一行,可以使用 `http://localhost:7474/` 在浏览器可视化访问。
```bash
-python -m src.api
+cd local_neo4j
+docker compose up -d
+```
+
+可以使用 `python test_neo4j.py` 来测试是否正常启动。使用 `docker compose down` 可停止服务。
+如果想要管理 neo4j,也可以使用 `docker ps` 查看容器 id,然后使用 `docker exec -it /bin/bash` 进入容器。
+如果想要删除数据库中的文件,可以进入容器并停止 neo4j 后,执行 `rm -rf /data/databases`。
+
+
+## 启动
+
+```bash
+python -m src.api
cd web
npm install
diff --git a/local_neo4j/docker-compose.yml b/local_neo4j/docker-compose.yml
new file mode 100644
index 00000000..aff1e9b2
--- /dev/null
+++ b/local_neo4j/docker-compose.yml
@@ -0,0 +1,17 @@
+version: '3.9'
+services:
+
+ neo4j:
+ image: neo4j:latest
+ volumes:
+ - ./conf:/var/lib/neo4j/conf
+ - ./import:/var/lib/neo4j/import
+ - ./plugins:/plugins
+ - ./data:/data
+ - ./logs:/var/lib/neo4j/logs
+ restart: always
+ ports:
+ - 7474:7474
+ - 7687:7687
+ environment:
+ - NEO4J_AUTH=neo4j/0123456789
diff --git a/local_neo4j/test_neo4j.py b/local_neo4j/test_neo4j.py
new file mode 100644
index 00000000..a293e4fc
--- /dev/null
+++ b/local_neo4j/test_neo4j.py
@@ -0,0 +1,34 @@
+from neo4j import GraphDatabase
+from neo4j.exceptions import ServiceUnavailable, AuthError
+
+def check_neo4j_status(uri="bolt://localhost:7687", username="neo4j", password="0123456789"):
+ """
+ 检查 Neo4j 数据库是否可以连接并正常工作。
+
+ 参数:
+ uri (str): Neo4j 的 URI,默认为 "bolt://localhost:7687"
+ username (str): 数据库用户名,默认为 "neo4j"
+ password (str): 数据库密码,默认为 "0123456789"
+
+ 返回:
+ str: "OK" 表示连接成功,"UNAVAILABLE" 表示服务不可用,"AUTH_FAILED" 表示认证失败。
+ """
+ try:
+ driver = GraphDatabase.driver(uri, auth=(username, password))
+ with driver.session() as session:
+ # 简单的查询来测试连接
+ result = session.run("RETURN 1")
+ if result.single()[0] == 1:
+ return "OK"
+ except ServiceUnavailable:
+ return "UNAVAILABLE"
+ except AuthError:
+ return "AUTH_FAILED"
+ finally:
+ # 确保关闭驱动
+ driver.close()
+
+# 测试函数
+status = check_neo4j_status()
+print(f"Neo4j status: {status}")
+
diff --git a/requirements.txt b/requirements.txt
new file mode 100644
index 00000000..0cbf80b3
--- /dev/null
+++ b/requirements.txt
@@ -0,0 +1,18 @@
+dashscope==1.20.5
+FlagEmbedding==1.2.11
+Flask==3.0.3
+Flask_Cors==4.0.1
+llama_index==0.11.1
+neo4j==5.23.1
+openai==1.42.0
+paddleocr==2.8.1
+pymilvus==2.4.5
+python-dotenv==1.0.1
+PyYAML==6.0.2
+qianfan==0.4.6
+torch==2.4.0
+tqdm==4.66.5
+zhipuai==2.1.4.20230814
+PyMuPDF
+llama-index-readers-file
+peft
\ No newline at end of file
diff --git a/src/config/__init__.py b/src/config/__init__.py
index b98038c4..50ae6ce7 100644
--- a/src/config/__init__.py
+++ b/src/config/__init__.py
@@ -38,12 +38,12 @@ class Config(SimpleConfig):
### >>> 默认配置
# 可以在 config/base.yaml 中覆盖
- self.add_item("mode", default="cli", des="运行模式", choices=["cli", "api"])
self.add_item("stream", default=True, des="是否开启流式输出")
self.add_item("save_dir", default="saves", des="保存目录")
# 功能选项
self.add_item("enable_reranker", default=False, des="是否开启重排序")
self.add_item("enable_knowledge_base", default=False, des="是否开启知识库")
+ self.add_item("enable_knowledge_graph", default=False, des="是否开启知识图谱")
self.add_item("enable_search_engine", default=False, des="是否开启搜索引擎")
# 模型配置
@@ -70,6 +70,13 @@ class Config(SimpleConfig):
"choices": choices
}
+ def __dict__(self):
+ blocklist = [
+ "_config_items",
+ "model_names",
+ ]
+ return {k: v for k, v in self.items() if k not in blocklist}
+
def handle_self(self):
### handle local model
model_root_dir = os.getenv("MODEL_ROOT_DIR", "pretrained_models")
@@ -98,7 +105,6 @@ class Config(SimpleConfig):
content = f.read()
if content:
local_config = json.loads(content)
- local_config.pop("_config_items")
self.update(local_config)
else:
print(f"{self.filename} is empty.")
@@ -108,7 +114,6 @@ class Config(SimpleConfig):
content = f.read()
if content:
local_config = yaml.safe_load(content)
- local_config.pop("_config_items")
self.update(local_config)
else:
print(f"{self.filename} is empty.")
diff --git a/src/core/database.py b/src/core/database.py
index 05b7b432..3966b6dc 100644
--- a/src/core/database.py
+++ b/src/core/database.py
@@ -8,46 +8,6 @@ from src.models.embedding import get_embedding_model
logger = setup_logger("DataBaseManager")
-class DataBaseLite:
- 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", {})
- self.files = kwargs.get("files", [])
- self.embed_model = kwargs.get("embed_model", None)
-
- def id2file(self, file_id):
- for f in self.files:
- if f["file_id"] == file_id:
- return f
- return None
-
- def update(self, metadata):
- self.metadata = metadata
-
- def to_dict(self):
- return {
- "name": self.name,
- "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,
- "dimension": self.dimension
- }
-
- def to_json(self):
- return json.dumps(self.to_dict(), ensure_ascii=False)
-
- def __str__(self):
- return self.to_json()
-
class DataBaseManager:
def __init__(self, config=None) -> None:
@@ -111,13 +71,16 @@ class DataBaseManager:
return {"databases": [db.to_dict() for db in self.data["databases"]]}
def get_graph(self):
- if self.config.enable_graph_base:
+ if self.config.enable_knowledge_graph:
self.data["graph"].update(self.graph_base.get_database_info("neo4j"))
return {"graph": self.data["graph"]}
else:
return {"message": "Graph base not enabled", "graph": {}}
def create_database(self, database_name, description, db_type, dimension):
+ from src.config import EMBED_MODEL_INFO
+ dimension = dimension or EMBED_MODEL_INFO[self.config.embed_model]["dimension"]
+
new_database = DataBaseLite(database_name,
description,
db_type,
@@ -134,7 +97,7 @@ class DataBaseManager:
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"}
+ return {"message": f"Embed model not match, cur: {self.config.embed_model}", "status": "failed"}
new_files = []
for file in files:
@@ -208,7 +171,6 @@ class DataBaseManager:
logger.error(f"File format not supported, only support {support_format}")
raise Exception(f"File format not supported, only support {support_format}")
-
def delete_file(self, db_id, file_id):
db = self.get_kb_by_id(db_id)
file_idx_to_delete = [idx for idx, f in enumerate(db.files) if f["file_id"] == file_id][0]
@@ -252,4 +214,45 @@ class DataBaseManager:
for db in self.data["databases"]:
if db.db_id == db_id:
return db
- return None
\ No newline at end of file
+ return None
+
+
+class DataBaseLite:
+ 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", {})
+ self.files = kwargs.get("files", [])
+ self.embed_model = kwargs.get("embed_model", None)
+
+ def id2file(self, file_id):
+ for f in self.files:
+ if f["file_id"] == file_id:
+ return f
+ return None
+
+ def update(self, metadata):
+ self.metadata = metadata
+
+ def to_dict(self):
+ return {
+ "name": self.name,
+ "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,
+ "dimension": self.dimension
+ }
+
+ def to_json(self):
+ return json.dumps(self.to_dict(), ensure_ascii=False)
+
+ def __str__(self):
+ return self.to_json()
\ No newline at end of file
diff --git a/src/core/graphbase.py b/src/core/graphbase.py
index 70d97255..664b4729 100644
--- a/src/core/graphbase.py
+++ b/src/core/graphbase.py
@@ -9,11 +9,12 @@ import warnings
from src.plugins import pdf2txt
from src.plugins.oneke import OneKE
+from src.utils import setup_logger
+
+logger = setup_logger("server-graphbase")
warnings.filterwarnings("ignore", category=UserWarning)
-
-
UIE_MODEL = None
class GraphDatabase:
@@ -36,6 +37,16 @@ class GraphDatabase:
"""关闭数据库连接"""
self.driver.close()
+ def get_sample_nodes(self, kgdb_name='neo4j', num=50):
+ """获取指定数据库的前 num 个节点信息"""
+ self.use_database(kgdb_name)
+ def query(tx, num):
+ result = tx.run("MATCH (n)-[r]->(m) RETURN n, r, m LIMIT $num", num=int(num))
+ return result.values()
+
+ with self.driver.session() as session:
+ return session.execute_read(query, num)
+
def create_graph_database(self, kgdb_name):
"""创建新的数据库,如果已存在则返回已有数据库的名称"""
with self.driver.session() as session:
@@ -116,21 +127,25 @@ class GraphDatabase:
MERGE (t:Entity {name: $t})
MERGE (h)-[r:RELATION {type: $r}]->(t)
""", h=entry['h'], t=entry['t'], r=entry['r'])
- def _create_vector_index(tx):
- index_name = "entity-embeddings"
+ def _create_vector_index(tx, dim):
+ index_name = "entityEmbeddings"
if not _index_exists(tx, index_name):
tx.run(f"""
CREATE VECTOR INDEX {index_name}
FOR (n: Entity) ON (n.embedding)
OPTIONS {{indexConfig: {{
- `vector.dimensions`: 1024,
+ `vector.dimensions`: {dim},
`vector.similarity_function`: 'cosine'
}} }};
""")
+
+ from src.config import EMBED_MODEL_INFO
+ embed_info = EMBED_MODEL_INFO[self.config.embed_model]
with self.driver.session() as session:
session.execute_write(_create_graph, triples)
- session.execute_write(_create_vector_index)
- for entry in triples:
+ session.execute_write(_create_vector_index, embed_info.dimension)
+ for i, entry in enumerate(triples):
+ logger.info(f"Adding entity {i+1}/{len(triples)}")
embedding_h = self.get_embedding(entry['h'])
session.execute_write(self.set_embedding, entry['h'], embedding_h)
@@ -148,37 +163,39 @@ class GraphDatabase:
triples = list(read_triples(file_path))
- def batch_create(tx, triples):
- query = """
- UNWIND $triples AS triple
- MERGE (a:Entity {name: triple.h})
- MERGE (b:Entity {name: triple.t})
- MERGE (a)-[r:RELATION {type: triple.r}]->(b)
- """
- tx.run(query, triples=triples)
+ self.txt_add_vector_entity(triples, kgdb_name)
- def batch_add_embeddings(tx, embeddings):
- query = """
- UNWIND $embeddings AS embedding
- MATCH (e:Entity {name: embedding.name})
- SET e.embedding = embedding.vector
- """
- tx.run(query, embeddings=embeddings)
-
- with self.driver.session() as session:
- session.execute_write(batch_create, triples)
-
- # 获取embedding并批量添加
- embeddings = []
- for triple in triples:
- h = triple['h']
- t = triple['t']
- embedding_h = self.get_embedding(h)
- embedding_t = self.get_embedding(t)
- embeddings.append({"name": h, "vector": embedding_h})
- embeddings.append({"name": t, "vector": embedding_t})
-
- session.execute_write(batch_add_embeddings, embeddings)
+ # def batch_create(tx, triples):
+ # query = """
+ # UNWIND $triples AS triple
+ # MERGE (a:Entity {name: triple.h})
+ # MERGE (b:Entity {name: triple.t})
+ # MERGE (a)-[r:RELATION {type: triple.r}]->(b)
+ # """
+ # tx.run(query, triples=triples)
+ #
+ # def batch_add_embeddings(tx, embeddings):
+ # query = """
+ # UNWIND $embeddings AS embedding
+ # MATCH (e:Entity {name: embedding.name})
+ # SET e.embedding = embedding.vector
+ # """
+ # tx.run(query, embeddings=embeddings)
+ #
+ # with self.driver.session() as session:
+ # session.execute_write(batch_create, triples)
+ #
+ # # 获取embedding并批量添加
+ # embeddings = []
+ # for triple in triples:
+ # h = triple['h']
+ # t = triple['t']
+ # embedding_h = self.get_embedding(h)
+ # embedding_t = self.get_embedding(t)
+ # embeddings.append({"name": h, "vector": embedding_h})
+ # embeddings.append({"name": t, "vector": embedding_t})
+ #
+ # session.execute_write(batch_add_embeddings, embeddings)
self.status = "open"
return kgdb_name
@@ -260,13 +277,21 @@ class GraphDatabase:
with self.driver.session() as session:
return session.execute_read(query, keyword, hops)
+ def query_node(self, entity_name, args):
+ # TODO 添加判断节点数量为 0 停止检索
+
+ if args.get("exact_match"):
+ raise NotImplemented("not implement for `exact_match`")
+ else:
+ return self.query_by_vector(entity_name, kgdb_name=args.get("kgdb_name"), hops=args.get("hops"))
+
def query_by_vector_tep(self, keyword, kgdb_name='neo4j'):
"""向量查询"""
self.use_database(kgdb_name)
def query(tx, text):
embedding = self.get_embedding(text)
result = tx.run("""
- CALL db.index.vector.queryNodes('entity-embeddings', 10, $embedding)
+ CALL db.index.vector.queryNodes('entityEmbeddings', 10, $embedding)
YIELD node AS similarEntity, score
RETURN similarEntity.name AS name, score
""", embedding=embedding)
@@ -277,7 +302,7 @@ class GraphDatabase:
with self.driver.session() as session:
return session.execute_read(query, keyword)
- def query_by_vector(self, entity_name, threshold=0.9,kgdb_name='neo4j', hops=2, num_of_res=2):
+ def query_by_vector(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, num_of_res=2):
self.use_database(kgdb_name)
result = self.query_by_vector_tep(entity_name)
querys = []
diff --git a/src/core/retriever.py b/src/core/retriever.py
index 2025f3a2..2071f3cb 100644
--- a/src/core/retriever.py
+++ b/src/core/retriever.py
@@ -83,7 +83,7 @@ class Retriever:
r["file"] = kb.id2file(r["entity"]["file_id"])
if self.config.enable_reranker:
- RERANK_THRESHOLD = 0.1
+ RERANK_THRESHOLD = 0.001
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)
@@ -124,7 +124,46 @@ class Retriever:
return entities
+ def foramt_general_results(self, results):
+ logger.debug(f"Formatting general results: {results}")
+ formatted_results = {"nodes": [], "edges": []}
+
+ for item in results:
+ relationship = item[1]
+ rel_id = relationship.element_id
+ nodes = relationship.nodes
+ if len(nodes) != 2:
+ continue
+
+ source, target = nodes
+
+ source_id = source.element_id
+ target_id = target.element_id
+ source_name = source._properties.get('name', 'unknown')
+ target_name = target._properties.get('name', 'unknown')
+
+ if source_id not in formatted_results["nodes"]:
+ formatted_results["nodes"].append({"id": source_id, "name": source_name})
+ if target_id not in formatted_results["nodes"]:
+ formatted_results["nodes"].append({"id": target_id, "name": target_name})
+
+ relationship_type = relationship._properties.get('type', 'unknown')
+ if relationship_type == 'unknown':
+ relationship_type = relationship.type
+
+ formatted_results["edges"].append({
+ "id": rel_id,
+ "type": relationship_type,
+ "source_id": source_id,
+ "target_id": target_id,
+ "source_name": source_name,
+ "target_name": target_name
+ })
+
+ return formatted_results
+
def format_query_results(self, results):
+ logger.debug(f"Formatting query results: {results}")
formatted_results = {"nodes": [], "edges": []}
node_dict = {}
diff --git a/src/models/embedding.py b/src/models/embedding.py
index e7c482d9..9eb81ee2 100644
--- a/src/models/embedding.py
+++ b/src/models/embedding.py
@@ -45,10 +45,11 @@ class ZhipuEmbedding:
self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:"
def predict(self, message):
-
data = []
for i in range(0, len(message), 10):
+ if len(message) > 10:
+ logger.info(f"Encoding {i} to {i+10} with {len(message)} messages")
group_msg = message[i:i+10]
response = self.client.embeddings.create(
model=self.model_info.default_path,
diff --git a/src/requirements.txt b/src/requirements.txt
deleted file mode 100644
index 29577253..00000000
--- a/src/requirements.txt
+++ /dev/null
@@ -1,6 +0,0 @@
-FlagEmbedding==1.2.10
-Flask==3.0.3
-Flask_Cors==4.0.1
-openai==1.35.10
-python-dotenv==1.0.1
-PyYAML==6.0.1
diff --git a/src/views/database_view.py b/src/views/database_view.py
index 43bf1b75..98adcfae 100644
--- a/src/views/database_view.py
+++ b/src/views/database_view.py
@@ -123,9 +123,20 @@ def get_graph_node():
return jsonify({'message': 'entity_name and kgdb_name are required'}), 400
logger.debug(f"Get graph node {entity_name} in {kgdb_name} with {hops} hops")
- result = startup.dbm.graph_base.query_by_vector(entity_name, kgdb_name=kgdb_name, hops=hops)
+ result = startup.dbm.graph_base.query_node(entity_name, request.args)
return jsonify({'result': startup.retriever.format_query_results(result), 'message': 'success'}), 200
+@db.route('/graph/nodes', methods=['GET'])
+def get_graph_nodes():
+ kgdb_name = request.args.get('kgdb_name')
+ num = request.args.get('num')
+ if not kgdb_name:
+ return jsonify({'message': 'kgdb_name is required'}), 400
+
+ logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes")
+ result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num)
+ return jsonify({'result': startup.retriever.foramt_general_results(result), 'message': 'success'}), 200
+
@db.route('/graph/add', methods=['POST'])
def add_graph_entity():
data = json.loads(request.data)
diff --git a/web/index.html b/web/index.html
index 251b07a4..e61e329a 100644
--- a/web/index.html
+++ b/web/index.html
@@ -1,5 +1,5 @@
-
+
diff --git a/web/package.json b/web/package.json
index 765b9d14..e95d555b 100644
--- a/web/package.json
+++ b/web/package.json
@@ -12,7 +12,7 @@
},
"dependencies": {
"@ant-design/icons-vue": "^6.1.0",
- "@antv/g6": "^5.0.9",
+ "@antv/g6": "^5.0.17",
"@vueuse/core": "^10.11.0",
"ant-design-vue": "^4.2.3",
"axios": "^1.3.4",
diff --git a/web/src/assets/base.css b/web/src/assets/base.css
index dbc51192..fcc4b5d6 100644
--- a/web/src/assets/base.css
+++ b/web/src/assets/base.css
@@ -11,6 +11,7 @@
--main-100: #ABE0F7;
--main-50: #CDF5FF;
--main-25: #E6FAFF;
+ --main-10: #F5FDFF;
--c-white: #ffffff;
--c-white-soft: #f8f8f8;
diff --git a/web/src/assets/main.css b/web/src/assets/main.css
index 9005a704..2deb3a7b 100644
--- a/web/src/assets/main.css
+++ b/web/src/assets/main.css
@@ -2,4 +2,16 @@
:root {
--header-height: 60px;
+}
+
+/* layout */
+
+.layout-container {
+ width: 100%;
+ padding: 0px 30px;
+ background-color: #FCFEFF;
+
+ h2 {
+ margin: 20px 0 10px 0;
+ }
}
\ No newline at end of file
diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue
index 93f51c8a..3b72d028 100644
--- a/web/src/components/ChatComponent.vue
+++ b/web/src/components/ChatComponent.vue
@@ -14,7 +14,7 @@
class="newchat nav-btn"
@click="$emit('newconv')"
>
- 新对话 {{ configStore.config?.model_name }}
+ 新对话:{{ configStore.config?.model_name }}