在检索测试里面支持 HyDE
This commit is contained in:
parent
07a77984a1
commit
4550635dc1
@ -73,15 +73,17 @@ class KnowledgeBase:
|
||||
def search(self, query, collection_name, limit=3):
|
||||
|
||||
query_vectors = self.embed_model.encode_queries([query])
|
||||
return self.search_by_vector(query_vectors[0], collection_name, limit)
|
||||
|
||||
def search_by_vector(self, vector, collection_name, limit=3):
|
||||
res = self.client.search(
|
||||
collection_name=collection_name, # target collection
|
||||
data=query_vectors, # query vectors
|
||||
data=[vector], # query vectors
|
||||
limit=limit, # number of returned entities
|
||||
output_fields=["text", "file_id"], # specifies fields to be returned
|
||||
)
|
||||
|
||||
return res[0] # 因为 query 只有一个
|
||||
return res[0]
|
||||
|
||||
def examples(self, collection_name, limit=20):
|
||||
res = self.client.query(
|
||||
|
||||
@ -15,9 +15,7 @@ class Retriever:
|
||||
|
||||
def retrieval(self, query, history, meta):
|
||||
|
||||
refs = {}
|
||||
refs["meta"] = meta
|
||||
refs["rewritten_query"] = self.rewrite_query(query, history, refs)
|
||||
refs = {"query": query, "history": history, "meta": meta}
|
||||
refs["entities"] = self.reco_entities(query, history, refs)
|
||||
refs["knowledge_base"] = self.query_knowledgebase(query, history, refs)
|
||||
refs["graph_base"] = self.query_graph(query, history, refs)
|
||||
@ -68,37 +66,47 @@ class Retriever:
|
||||
|
||||
def query_knowledgebase(self, query, history, refs):
|
||||
"""查询知识库"""
|
||||
query = refs.get("rewritten_query", query)
|
||||
rw_query = self.rewrite_query(query, history, refs)
|
||||
|
||||
kb_res = []
|
||||
final_res = []
|
||||
if refs["meta"].get("db_name") and self.config.enable_knowledge_base:
|
||||
|
||||
db_name = refs["meta"]["db_name"]
|
||||
kb = self.dbm.metaname2db[refs["meta"]["db_name"]]
|
||||
kb = self.dbm.metaname2db[db_name]
|
||||
limit = refs["meta"].get("queryCount", 10)
|
||||
kb_res = self.dbm.knowledge_base.search(query, db_name, limit=limit)
|
||||
|
||||
kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=limit)
|
||||
for r in kb_res:
|
||||
r["file"] = kb.id2file(r["entity"]["file_id"])
|
||||
|
||||
if self.config.enable_reranker:
|
||||
RERANK_THRESHOLD = 0.1
|
||||
for r in kb_res:
|
||||
r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True)
|
||||
kb_res.sort(key=lambda x: x["rerank_score"], reverse=True)
|
||||
final_res = [_res for _res in kb_res if _res["rerank_score"] > 0.1]
|
||||
final_res = [_res for _res in kb_res if _res["rerank_score"] > RERANK_THRESHOLD]
|
||||
|
||||
else:
|
||||
final_res = kb_res[:5]
|
||||
|
||||
return {"results": final_res, "all_results": kb_res}
|
||||
return {"results": final_res, "all_results": kb_res, "rw_query": rw_query}
|
||||
|
||||
def rewrite_query(self, query, history, refs):
|
||||
"""重写查询"""
|
||||
if refs["meta"].get("rewrite_query") is None or history == []:
|
||||
rewrite_query_span = refs["meta"].get("rewrite_query", None)
|
||||
if rewrite_query_span is None or rewrite_query_span == "OFF":
|
||||
rewritten_query = query
|
||||
else:
|
||||
from src.utils.prompts import rewritten_query_prompt_template
|
||||
rewritten_query_prompt = rewritten_query_prompt_template.format(history=[entry['content'] for entry in history if entry['role'] == 'user'], query=query)
|
||||
history_query = [entry['content'] for entry in history if entry['role'] == 'user'] if history else ""
|
||||
rewritten_query_prompt = rewritten_query_prompt_template.format(history=history_query, query=query)
|
||||
rewritten_query = self.model.predict(rewritten_query_prompt).content
|
||||
|
||||
if rewrite_query_span == "HyDE":
|
||||
hy_doc = self.model.predict(rewritten_query).content
|
||||
rewritten_query = f"{rewritten_query} {hy_doc}"
|
||||
|
||||
return rewritten_query
|
||||
|
||||
def reco_entities(self, query, history, refs):
|
||||
|
||||
@ -34,7 +34,12 @@
|
||||
</div>
|
||||
<div class="params-item">
|
||||
<p>过滤低质量:</p>
|
||||
<a-switch v-model:checked="meta.filter" size="small" />
|
||||
<a-switch v-model:checked="meta.filter" />
|
||||
</div>
|
||||
<div class="params-item">
|
||||
<p>重写查询</p>
|
||||
<a-segmented v-model:value="meta.rewrite_query" :options="['OFF', 'ON', 'HyDE']" />
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@ -122,7 +127,7 @@
|
||||
:auto-size="{ minRows: 2, maxRows: 10 }"
|
||||
/>
|
||||
<!-- :loading="state.searchLoading" -->
|
||||
<a-button @click="onQuery" :disabled="queryText.length == 0">
|
||||
<a-button @click="onQuery" :disabled="queryText.length == 0" :loading="state.searchLoading">
|
||||
<SearchOutlined v-if="!state.searchLoading"/>检索
|
||||
</a-button>
|
||||
</div>
|
||||
@ -183,6 +188,7 @@ const state = reactive({
|
||||
const meta = reactive({
|
||||
queryCount: 10,
|
||||
filter: false,
|
||||
rewrite_query: 'OFF',
|
||||
});
|
||||
|
||||
const onQuery = () => {
|
||||
@ -507,9 +513,8 @@ onMounted(() => {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
padding: 20px;
|
||||
background: var(--main-light-4);
|
||||
border-radius: 8px;
|
||||
margin: 20px;
|
||||
margin: 20px 0;
|
||||
box-sizing: border-box;
|
||||
gap: 12px;
|
||||
|
||||
@ -523,6 +528,10 @@ onMounted(() => {
|
||||
margin: 0;
|
||||
}
|
||||
}
|
||||
.params-item.col {
|
||||
flex-direction: column;
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user