From b432c1bda9b3b6eb5ac2697565566d822d77f375 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Mon, 31 Mar 2025 18:32:28 +0800 Subject: [PATCH] =?UTF-8?q?=E6=A3=80=E7=B4=A2=E5=89=8D=E6=A3=80=E6=9F=A5?= =?UTF-8?q?=E6=98=AF=E5=90=A6=E5=AD=98=E5=9C=A8=20entityEmbeddings=20?= =?UTF-8?q?=E8=BF=99=E4=B8=AA=E7=B4=A2=E5=BC=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docker/docker-compose.dev.yml | 1 + src/core/graphbase.py | 23 ++++++++++++++++++++--- src/core/retriever.py | 2 ++ 3 files changed, 23 insertions(+), 3 deletions(-) diff --git a/docker/docker-compose.dev.yml b/docker/docker-compose.dev.yml index 0232dc2d..dfd09a6f 100644 --- a/docker/docker-compose.dev.yml +++ b/docker/docker-compose.dev.yml @@ -60,6 +60,7 @@ services: - NEO4J_AUTH=neo4j/0123456789 - NEO4J_server_bolt_listen__address=0.0.0.0:7687 - NEO4J_server_http_listen__address=0.0.0.0:7474 + - ENTITY_EMBEDDING=true networks: - app-network diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 75bf1419..020ef64a 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -218,7 +218,19 @@ class GraphDatabase: def query_by_vector(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, num_of_res=5): self.use_database(kgdb_name) + def _index_exists(tx, index_name): + """检查索引是否存在""" + result = tx.run("SHOW INDEXES") + for record in result: + if record["name"] == index_name: + return True + return False + def query(tx, text): + # 首先检查索引是否存在 + if not _index_exists(tx, "entityEmbeddings"): + raise Exception("向量索引不存在,请先创建索引") + embedding = self.get_embedding(text) result = tx.run(""" CALL db.index.vector.queryNodes('entityEmbeddings', 10, $embedding) @@ -227,9 +239,14 @@ class GraphDatabase: """, embedding=embedding) return result.values() - with self.driver.session() as session: - results = session.execute_read(query, entity_name) - + try: + with self.driver.session() as session: + results = session.execute_read(query, entity_name) + except Exception as e: + if "向量索引不存在" in str(e): + logger.error(f"向量索引不存在,请先创建索引: {e}, {traceback.format_exc()}") + return [] + raise e # 筛选出分数高于阈值的实体 qualified_entities = [result[0] for result in results[:num_of_res] if result[1] > threshold] diff --git a/src/core/retriever.py b/src/core/retriever.py index 3181a870..ee52d915 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -76,6 +76,8 @@ class Retriever: results = [] if refs["meta"].get("use_graph") and config.enable_knowledge_base: for entity in refs["entities"]: + if entity == "": + continue result = graph_base.query_by_vector(entity) if result != []: results.extend(result)