From 38e3f27899e4739dd4f56847d81a29d837e953da Mon Sep 17 00:00:00 2001 From: Soulter <905617992@qq.com> Date: Fri, 24 Oct 2025 15:06:07 +0800 Subject: [PATCH] feat: update knowledge base retrieval configuration and UI adjustments --- astrbot/core/config/default.py | 9 ++++++++- astrbot/core/knowledge_base/kb_mgr.py | 2 ++ astrbot/core/knowledge_base/retrieval/manager.py | 6 ++++-- .../knowledge_base/retrieval/sparse_retriever.py | 3 ++- astrbot/core/pipeline/process_stage/utils.py | 2 ++ .../views/knowledge-base/components/SettingsTab.vue | 12 ++++++------ 6 files changed, 24 insertions(+), 10 deletions(-) diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index 36ac616ae..497f03eae 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -137,7 +137,8 @@ DEFAULT_CONFIG = { "default_kb_collection": "", # 默认知识库名称, 已经过时 "plugin_set": ["*"], # "*" 表示使用所有可用的插件, 空列表表示不使用任何插件 "kb_names": [], # 默认知识库名称列表 - "kb_final_top_k": 5, # 知识库检索的最终返回结果数量 + "kb_fusion_top_k": 20, # 知识库检索融合阶段返回结果数量 + "kb_final_top_k": 5, # 知识库检索最终返回结果数量 } @@ -2003,6 +2004,7 @@ CONFIG_METADATA_2 = { "type": "string", }, "kb_names": {"type": "list", "items": {"type": "string"}}, + "kb_fusion_top_k": {"type": "int", "default": 20}, "kb_final_top_k": {"type": "int", "default": 5}, }, }, @@ -2089,6 +2091,11 @@ CONFIG_METADATA_3 = { "_special": "select_knowledgebase", "hint": "支持多选", }, + "kb_fusion_top_k": { + "description": "融合检索结果数", + "type": "int", + "hint": "多个知识库检索结果融合后的返回结果数量", + }, "kb_final_top_k": { "description": "最终返回结果数", "type": "int", diff --git a/astrbot/core/knowledge_base/kb_mgr.py b/astrbot/core/knowledge_base/kb_mgr.py index 47a48af64..b079b73e4 100644 --- a/astrbot/core/knowledge_base/kb_mgr.py +++ b/astrbot/core/knowledge_base/kb_mgr.py @@ -212,6 +212,7 @@ class KnowledgeBaseManager: self, query: str, kb_names: list[str], + top_k_fusion: int = 20, top_m_final: int = 5, ) -> dict | None: """从指定知识库中检索相关内容""" @@ -229,6 +230,7 @@ class KnowledgeBaseManager: query=query, kb_ids=kb_ids, kb_id_helper_map=kb_id_helper_map, + top_k_fusion=top_k_fusion, top_m_final=top_m_final, ) if not results: diff --git a/astrbot/core/knowledge_base/retrieval/manager.py b/astrbot/core/knowledge_base/retrieval/manager.py index 7e90cf2f6..00d1649b0 100644 --- a/astrbot/core/knowledge_base/retrieval/manager.py +++ b/astrbot/core/knowledge_base/retrieval/manager.py @@ -61,6 +61,7 @@ class RetrievalManager: query: str, kb_ids: List[str], kb_id_helper_map: dict[str, KBHelper], + top_k_fusion: int = 20, top_m_final: int = 5, ) -> List[RetrievalResult]: """混合检索 @@ -120,7 +121,7 @@ class RetrievalManager: fused_results = await self.rank_fusion.fuse( dense_results=dense_results, sparse_results=sparse_results, - top_k=kb_options.get("top_k_fusion", 20), + top_k=top_k_fusion, ) # 4. 转换为 RetrievalResult (获取元数据) @@ -209,7 +210,8 @@ class RetrievalManager: # 按相似度排序并返回 top_k all_results.sort(key=lambda x: x.similarity, reverse=True) - return all_results[: len(all_results) // len(kb_ids)] + # return all_results[: len(all_results) // len(kb_ids)] + return all_results async def _rerank( self, diff --git a/astrbot/core/knowledge_base/retrieval/sparse_retriever.py b/astrbot/core/knowledge_base/retrieval/sparse_retriever.py index ca8285be5..4275c3ef4 100644 --- a/astrbot/core/knowledge_base/retrieval/sparse_retriever.py +++ b/astrbot/core/knowledge_base/retrieval/sparse_retriever.py @@ -124,4 +124,5 @@ class SparseRetriever: ) results.sort(key=lambda x: x.score, reverse=True) - return results[: len(results) // len(kb_ids)] + # return results[: len(results) // len(kb_ids)] + return results diff --git a/astrbot/core/pipeline/process_stage/utils.py b/astrbot/core/pipeline/process_stage/utils.py index 8fe34e3ee..7ba3757c0 100644 --- a/astrbot/core/pipeline/process_stage/utils.py +++ b/astrbot/core/pipeline/process_stage/utils.py @@ -16,6 +16,7 @@ async def inject_kb_context( """ kb_mgr = p_ctx.plugin_manager.context.kb_manager kb_names = p_ctx.astrbot_config.get("kb_names", []) + top_k_fusion = p_ctx.astrbot_config.get("kb_fusion_top_k", 20) top_k = p_ctx.astrbot_config.get("kb_final_top_k", 5) if not kb_names: @@ -24,6 +25,7 @@ async def inject_kb_context( kb_context = await kb_mgr.retrieve( query=req.prompt, kb_names=kb_names, + top_k_fusion=top_k_fusion, top_m_final=top_k, ) if not kb_context: diff --git a/dashboard/src/views/knowledge-base/components/SettingsTab.vue b/dashboard/src/views/knowledge-base/components/SettingsTab.vue index 0d468ba98..fb3911aaa 100644 --- a/dashboard/src/views/knowledge-base/components/SettingsTab.vue +++ b/dashboard/src/views/knowledge-base/components/SettingsTab.vue @@ -34,7 +34,7 @@