优化重写查询

This commit is contained in:
Wenjie Zhang 2025-04-11 11:53:37 +08:00
parent 3a67a930d9
commit f11f719b4e
3 changed files with 65 additions and 3 deletions

48
src/core/operators.py Normal file
View File

@ -0,0 +1,48 @@
"""
这里面存放是 RAG 相关的一些组件
"""
from src.utils import prompts
class BaseOperator:
"""
基类
"""
template = None
def __init__(self):
pass
def call(self, **kwargs):
pass
def __call__(self, **kwargs):
"""
所有 RAG 相关组件的调用接口
"""
return self.call(**kwargs)
class HyDEOperator(BaseOperator):
"""
HyDE 重写查询
"""
template = prompts.HYDE_PROMPT_TEMPLATE
def __init__(self):
super().__init__()
@classmethod
def call(cls, model_callable, query, context_str, **kwargs):
"""
重写查询
Args:
model_callable: 模型调用函数
query: 查询
context_str: 上下文
"""
prompt = cls.template.format(query=query, context_str=context_str)
response = model_callable(prompt)
return response

View File

@ -2,6 +2,8 @@ from src import config, knowledge_base, graph_base
from src.models.rerank_model import get_reranker
from src.utils.logging_config import logger
from src.models import select_model
from src.core.operators import HyDEOperator
class Retriever:
@ -151,8 +153,8 @@ class Retriever:
rewritten_query = model.predict(rewritten_query_prompt).content
if rewrite_query_span == "hyde":
hy_doc = model.predict(rewritten_query).content
rewritten_query = f"{rewritten_query} {hy_doc}"
res = HyDEOperator.call(model_callable=model.predict, query=query, context_str=history_query)
rewritten_query = res.content
return rewritten_query

View File

@ -56,4 +56,16 @@ keywords_prompt_template = """
返回的实体使用<->隔开关键词1<->关键词<->关键词3
不要改变关键词的语言
<文本>{text}</文本>
"""
"""
HYDE_PROMPT_TEMPLATE = (
"Please write a passage to answer the question\n"
"Try to include as many key details as possible.\n"
"\n"
"\n"
"{context_str}\n"
"\n"
"{query}\n"
"\n"
'Passage:\n'
)