diff --git a/.gitignore b/.gitignore index 990ebc34..f5509855 100644 --- a/.gitignore +++ b/.gitignore @@ -19,6 +19,7 @@ log logs *.log.* *.db +*.lock ### IDE .vscode diff --git a/src/api.py b/src/api.py index 30e47cb6..e2afe7fb 100644 --- a/src/api.py +++ b/src/api.py @@ -50,7 +50,7 @@ def chat(): if config.enable_knowledge_base: kb_res = pre_retrieval.search(query) if kb_res: - kb_res = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res[0]]) + kb_res = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res]) kb_res = f"知识库信息: {kb_res}" external += kb_res diff --git a/src/cli.py b/src/cli.py index 39257474..351c03d4 100644 --- a/src/cli.py +++ b/src/cli.py @@ -1,7 +1,7 @@ import os from dotenv import load_dotenv from core import HistoryManager -from core import PreRetrieval +from core import PreRetrieval, Retriever from config import Config from models import select_model @@ -11,30 +11,23 @@ load_dotenv() if __name__ == "__main__": config = Config("config/base.yaml") model = select_model(config) - pre_retrieval = PreRetrieval(config) + retriever = Retriever(config) # pre_retrieval.add_file("/home/zwj/workspace/ProjectAthena/src/data/file/鉴定工作报告、技术报告-0708.pdf") print(f"[{config.model_provider}:{config.get('model_name', 'default')}] Type 'exit' to quit") history_manager = HistoryManager() while True: - message = input("\nUser: ") - if message == "exit": + query = input("\nUser: ") + if query == "exit": break - external = "" + # 检索结果 + refs = retriever(query) + # 重新构建用户的输入 + query = retriever.construct_query(query, refs) - if config.enable_knowledge_base: - kb_res = pre_retrieval.search(message) - if kb_res: - kb_res = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res[0]]) - kb_res = f"知识库信息: {kb_res}" - external += kb_res - - if len(external) > 0: - message = f"以下是参考资料:\n\n\n {external} 请根据前面的知识回答:{message}" - - messages = history_manager.add_user(message) + messages = history_manager.add_user(query) response = model.predict(messages, stream=config.stream) if config.stream: diff --git a/src/core/__init__.py b/src/core/__init__.py index 71c46273..0338e6fe 100644 --- a/src/core/__init__.py +++ b/src/core/__init__.py @@ -1,2 +1,3 @@ from .history import * -from .preretrieval import * \ No newline at end of file +from .preretrieval import * +from .retriever import * \ No newline at end of file diff --git a/src/core/preretrieval.py b/src/core/preretrieval.py index 3ac793de..b3251505 100644 --- a/src/core/preretrieval.py +++ b/src/core/preretrieval.py @@ -88,7 +88,7 @@ class PreRetrieval: output_fields=["text", "subject"], # specifies fields to be returned ) - return res + return res[0] # 因为 query 只有一个 def read_text(self, file): support_format = [".pdf", ".txt", "*.md"] diff --git a/src/core/retrieval.py b/src/core/retrieval.py deleted file mode 100644 index e69de29b..00000000 diff --git a/src/core/retriever.py b/src/core/retriever.py new file mode 100644 index 00000000..a995ba38 --- /dev/null +++ b/src/core/retriever.py @@ -0,0 +1,50 @@ +from core import PreRetrieval + +class Retriever: + + def __init__(self, config): + self.config = config + self.pre_retrieval = PreRetrieval(config) + + def retrieval(self, query): + + refs = {} + + # TODO: 查询分类、查询重写、查询分解、查询伪文档生成(HyDE) + + if self.config.enable_knowledge_base: + refs["knowledge_base"] = self.pre_retrieval.search(query) + + return refs + + def construct_query(self, query, refs): + # TODO:Reranking + + if len(refs) == 0: + return query + + external = "" + + kb_res = refs.get("knowledge_base") + if kb_res: + kb_text = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res]) + external += f"知识库信息: \n\n{kb_text}" + + if len(external) > 0: + query = f"以下是参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答:{query}" + + return query + + def query_classification(self, query): + """判断是否需要查询 + - 对于完全基于用户给定信息的任务,称之为“足够”“sufficient”,不需要检索; + - 否则,称之为“不足”“insufficient”,可能需要检索, + """ + raise NotImplementedError + + def rewrite_query(self, query): + """重写查询""" + raise NotImplementedError + + def __call__(self, query): + return self.retrieval(query) \ No newline at end of file