update retriever

This commit is contained in:
Wenjie Zhang 2024-07-10 12:56:53 +08:00
parent a9124a1881
commit 39bd774159
7 changed files with 64 additions and 19 deletions

1
.gitignore vendored
View File

@ -19,6 +19,7 @@ log
logs
*.log.*
*.db
*.lock
### IDE
.vscode

View File

@ -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

View File

@ -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:

View File

@ -1,2 +1,3 @@
from .history import *
from .preretrieval import *
from .preretrieval import *
from .retriever import *

View File

@ -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"]

View File

50
src/core/retriever.py Normal file
View File

@ -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):
# TODOReranking
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)