diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 97f8855c..b4ed57b0 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -268,18 +268,18 @@ class GraphDatabase: return self.query_by_vector(entity_name=entity_name, **kwargs) def query_by_vector(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, num_of_res=5): - result = self.query_by_vector_tep(entity_name=entity_name) - querys = [] - for i in range(num_of_res): - if result[i][1] > threshold: - querys.append(result[i][0]) - else: - break - ans = [] - for query in querys: - tep = self.query_specific_entity(entity_name=query, hops=hops) # 这里是只获取第一个 TODO: 优化 - ans.extend(tep) - return ans + results = self.query_by_vector_tep(entity_name=entity_name) + + # 筛选出分数高于阈值的实体 + qualified_entities = [result[0] for result in results[:num_of_res] if result[1] > threshold] + + # 对每个合格的实体进行查询 + all_query_results = [] + for entity in qualified_entities: + query_result = self.query_specific_entity(entity_name=entity, hops=hops, kgdb_name=kgdb_name) + all_query_results.extend(query_result) + + return all_query_results def query_by_vector_tep(self, entity_name, kgdb_name='neo4j'): """向量查询""" diff --git a/src/core/retriever.py b/src/core/retriever.py index 52c4f5e7..788b4107 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -25,27 +25,28 @@ class Retriever: return refs def construct_query(self, query, refs, meta): - if len(refs) == 0: + if not refs or len(refs) == 0: return query - external = "" + external_parts = [] # 解析知识库的结果 - kb_res = refs.get("knowledge_base").get("results", []) - if len(kb_res) > 0: - kb_text = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res]) - external += f"知识库信息: \n\n{kb_text}" + kb_res = refs.get("knowledge_base", {}).get("results", []) + if kb_res: + kb_text = "\n".join(f"{r['id']}: {r['entity']['text']}" for r in kb_res) + external_parts.extend(["知识库信息:", kb_text]) # 解析图数据库的结果 - db_res = refs.get("graph_base").get("results", []) - if len(db_res["nodes"]) > 0: + db_res = refs.get("graph_base", {}).get("results", {}) + if db_res.get("nodes") and len(db_res["nodes"]) > 0: db_text = "\n".join( - [f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res["edges"]] + [f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res.get("edges", [])] ) - external += f"图数据库信息: \n\n{db_text}" + external_parts.extend(["图数据库信息:", db_text]) # 构造查询 - if len(external) > 0: + if external_parts and len(external_parts) > 0: + external = "\n\n".join(external_parts) query = f"参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答问题。\n\n问题:{query}\n\n回答:" return query @@ -73,7 +74,9 @@ class Retriever: kb_res = [] final_res = [] - if not refs["meta"].get("db_name") or not self.config.enable_knowledge_base: + + db_name = refs["meta"].get("db_name") + if not db_name or not self.config.enable_knowledge_base: return { "results": final_res, "all_results": kb_res, @@ -83,21 +86,21 @@ class Retriever: rw_query = self.rewrite_query(query, history, refs) - db_name = refs["meta"]["db_name"] kb = self.dbm.metaname2db[db_name] logger.debug(f"{refs['meta']=}") - max_query_count = refs["meta"].get("maxQueryCount", 10) - rerank_threshold = refs["meta"].get("rerankThreshold", 0.1) - distance_threshold = refs["meta"].get("distanceThreshold", 0) - top_k = refs["meta"].get("topK", 5) + meta = refs["meta"] + max_query_count = meta.get("maxQueryCount", 10) + rerank_threshold = meta.get("rerankThreshold", 0.1) + distance_threshold = meta.get("distanceThreshold", 0) + top_k = meta.get("topK", 5) all_kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=max_query_count) for r in all_kb_res: r["file"] = kb.id2file(r["entity"]["file_id"]) # use distance threshold to filter results - if refs["meta"].get("mode") == "search": + if meta.get("mode") == "search": kb_res = all_kb_res else: kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold] @@ -136,11 +139,12 @@ class Retriever: entities = [] if refs["meta"].get("use_graph"): - from src.utils.prompts import entity_extraction_prompt_template + from src.utils.prompts import entity_extraction_prompt_template as entity_template + from src.utils.prompts import keywords_prompt_template as entity_template - entity_extraction_prompt = entity_extraction_prompt_template.format(text=query) - 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)] + entity_extraction_prompt = entity_template.format(text=query) + 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)] return entities @@ -202,7 +206,7 @@ class Retriever: node_dict = {} for item in results: - if not isinstance(item[1], list) or len(item[1]) == 0: + if not isinstance(item[1], list) or not item[1]: continue relationship = item[1][0] @@ -213,10 +217,7 @@ class Retriever: if node_info is None or edge_info is None: continue - for node in node_info: - if node["id"] not in node_dict: - node_dict[node["id"]] = node - + node_dict.update({node["id"]: node for node in node_info}) formatted_results["edges"].append(edge_info) formatted_results["nodes"] = list(node_dict.values()) diff --git a/src/utils/prompts.py b/src/utils/prompts.py index eeb0ead1..208d067d 100644 --- a/src/utils/prompts.py +++ b/src/utils/prompts.py @@ -30,4 +30,12 @@ entity_extraction_prompt_template = """ 3.只返回实体,不得返回其他任何内容。 4.返回的实体用逗号隔开<内容要求> <文本>{text} +""" + +keywords_prompt_template = """ +你是用来辅助查询的助手,请对以下文本进行关键词提取,返回提取出的关键词。 +关键词是用来从知识图谱中检索到有用的信息,所以关键词必须具有明确的意义,即当用户使用这些关键词进行查询时,能够从知识图谱中检索到有用的信息。 +返回的实体使用<->隔开。如:关键词1<->关键词<->关键词3 +不要改变关键词的语言 +<文本>{text} """ \ No newline at end of file diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index b5dc9200..9773e86a 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -572,19 +572,19 @@ watch( gap: 10px; .opt__button { - background-color: white; - color: #222; - padding: 6px 1rem; - border-radius: 1rem; + background-color: #f2f5f5; + color: #333; + padding: .5rem 1.5rem; + border-radius: 2rem; cursor: pointer; // border: 2px solid var(--main-light-4); transition: background-color 0.3s; - box-shadow: 0px 0px 10px 4px var(--main-light-4); + // box-shadow: 0px 0px 10px 2px var(--main-light-4); &:hover { - background-color: #fcfcfc; - box-shadow: 0px 0px 10px 1px rgba(0, 0, 0, 0.1); + background-color: #f0f1f1; + // box-shadow: 0px 0px 10px 1px rgba(0, 0, 0, 0.1); } } } @@ -672,19 +672,19 @@ watch( max-width: 900px; margin: 0 auto; align-items: flex-end; - padding: 0.5rem; - box-shadow: rgba(42, 60, 79, 0.1) 0px 6px 10px 0px; - border: 1px solid #E5E5E5; - border-radius: 2rem; - background: #fafafa; + padding: 0.25rem 0.5rem; + // box-shadow: rgba(42, 60, 79, 0.1) 0px 6px 10px 0px; + border: 2px solid #E5E5E5; + border-radius: 1rem; + background: #fcfdfd; transition: background 0.3s, box-shadow 0.3s; &:focus-within { - border: 1px solid #BABABA; + border: 2px solid var(--main-500); background: white; - box-shadow: rgba(42, 60, 79, 0.1) 0px 6px 10px 0px; + // box-shadow: rgb(42 60 79 / 5%) 0px 4px 10px 0px; } - .user-input { + textarea.user-input { flex: 1; height: 40px; padding: 0.5rem 0.5rem; @@ -696,7 +696,7 @@ watch( font-size: 16px; font-variation-settings: 'wght' 400, 'opsz' 10.5; outline: none; - + resize: none; &:focus { outline: none; box-shadow: none; @@ -706,27 +706,27 @@ watch( outline: none; } } - - .send-btn { - border: none; - background: transparent; - cursor: pointer; - font-weight: 500; - padding: 0.5rem 1rem; - border-radius: 1rem; - transition: background-color 0.3s; - - &:hover { - background-color: var(--main-light-3); - } - - &:disabled { - cursor: not-allowed; - background-color: #DCDCDC; - } - } } + button.ant-btn-icon-only { + font-size: 1.25rem; + cursor: pointer; + background-color: transparent; + border: none; + transition: color 0.3s; + box-shadow: none; + color: var(--main-700);; + padding: 0; + + &:hover { + color: var(--c-text-dark-1); + } + + &:disabled { + color: #ccc; + cursor: not-allowed; + } + } .note { width: 100%; font-size: small; @@ -745,27 +745,6 @@ watch( cursor: pointer; } -.ant-btn-icon-only { - font-size: 16px; - cursor: pointer; - background-color: transparent; - border: none; - height: 2.5rem; - background-color: var(--main-color); - border-radius: 3rem; - color: white; - transition: background-color 0.3s; - - &:hover { - color: white; - } -} - -button:disabled { - background: #E0E0E0; - cursor: not-allowed; -} - .chat::-webkit-scrollbar { diff --git a/web/src/components/GraphContainer.vue b/web/src/components/GraphContainer.vue new file mode 100644 index 00000000..30324c20 --- /dev/null +++ b/web/src/components/GraphContainer.vue @@ -0,0 +1,95 @@ + + + + + \ No newline at end of file diff --git a/web/src/components/RefsComponent.vue b/web/src/components/RefsComponent.vue index ef2bbd0e..22c60354 100644 --- a/web/src/components/RefsComponent.vue +++ b/web/src/components/RefsComponent.vue @@ -1,17 +1,26 @@ - + +.results-list { + .result-item { + border-bottom: 1px solid #f0f0f0; + padding: 16px 0; + + &:last-child { + border-bottom: none; + } + } + + .result-meta { + margin-bottom: 12px; + } +} + \ No newline at end of file diff --git a/web/src/layouts/AppLayout.vue b/web/src/layouts/AppLayout.vue index a1581300..e27caa1c 100644 --- a/web/src/layouts/AppLayout.vue +++ b/web/src/layouts/AppLayout.vue @@ -58,8 +58,18 @@ console.log(route)