update retriever
This commit is contained in:
parent
a9124a1881
commit
39bd774159
1
.gitignore
vendored
1
.gitignore
vendored
@ -19,6 +19,7 @@ log
|
|||||||
logs
|
logs
|
||||||
*.log.*
|
*.log.*
|
||||||
*.db
|
*.db
|
||||||
|
*.lock
|
||||||
|
|
||||||
### IDE
|
### IDE
|
||||||
.vscode
|
.vscode
|
||||||
|
|||||||
@ -50,7 +50,7 @@ def chat():
|
|||||||
if config.enable_knowledge_base:
|
if config.enable_knowledge_base:
|
||||||
kb_res = pre_retrieval.search(query)
|
kb_res = pre_retrieval.search(query)
|
||||||
if kb_res:
|
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}"
|
kb_res = f"知识库信息: {kb_res}"
|
||||||
external += kb_res
|
external += kb_res
|
||||||
|
|
||||||
|
|||||||
25
src/cli.py
25
src/cli.py
@ -1,7 +1,7 @@
|
|||||||
import os
|
import os
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
from core import HistoryManager
|
from core import HistoryManager
|
||||||
from core import PreRetrieval
|
from core import PreRetrieval, Retriever
|
||||||
from config import Config
|
from config import Config
|
||||||
from models import select_model
|
from models import select_model
|
||||||
|
|
||||||
@ -11,30 +11,23 @@ load_dotenv()
|
|||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
config = Config("config/base.yaml")
|
config = Config("config/base.yaml")
|
||||||
model = select_model(config)
|
model = select_model(config)
|
||||||
pre_retrieval = PreRetrieval(config)
|
retriever = Retriever(config)
|
||||||
# pre_retrieval.add_file("/home/zwj/workspace/ProjectAthena/src/data/file/鉴定工作报告、技术报告-0708.pdf")
|
# 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")
|
print(f"[{config.model_provider}:{config.get('model_name', 'default')}] Type 'exit' to quit")
|
||||||
|
|
||||||
history_manager = HistoryManager()
|
history_manager = HistoryManager()
|
||||||
while True:
|
while True:
|
||||||
message = input("\nUser: ")
|
query = input("\nUser: ")
|
||||||
if message == "exit":
|
if query == "exit":
|
||||||
break
|
break
|
||||||
|
|
||||||
external = ""
|
# 检索结果
|
||||||
|
refs = retriever(query)
|
||||||
|
# 重新构建用户的输入
|
||||||
|
query = retriever.construct_query(query, refs)
|
||||||
|
|
||||||
if config.enable_knowledge_base:
|
messages = history_manager.add_user(query)
|
||||||
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)
|
|
||||||
response = model.predict(messages, stream=config.stream)
|
response = model.predict(messages, stream=config.stream)
|
||||||
|
|
||||||
if config.stream:
|
if config.stream:
|
||||||
|
|||||||
@ -1,2 +1,3 @@
|
|||||||
from .history import *
|
from .history import *
|
||||||
from .preretrieval import *
|
from .preretrieval import *
|
||||||
|
from .retriever import *
|
||||||
@ -88,7 +88,7 @@ class PreRetrieval:
|
|||||||
output_fields=["text", "subject"], # specifies fields to be returned
|
output_fields=["text", "subject"], # specifies fields to be returned
|
||||||
)
|
)
|
||||||
|
|
||||||
return res
|
return res[0] # 因为 query 只有一个
|
||||||
|
|
||||||
def read_text(self, file):
|
def read_text(self, file):
|
||||||
support_format = [".pdf", ".txt", "*.md"]
|
support_format = [".pdf", ".txt", "*.md"]
|
||||||
|
|||||||
50
src/core/retriever.py
Normal file
50
src/core/retriever.py
Normal 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):
|
||||||
|
# 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)
|
||||||
Loading…
Reference in New Issue
Block a user