kgrag/app_pkg/examples/search.py
2026-07-29 18:10:19 +08:00

158 lines
6.7 KiB
Python

from pymilvus import Collection, db, connections,
import numpy as np
from typing import List
from openai import OpenAI
import os
openai_embedding_api_base = os.getenv("OPENAI_API_EMBEDDING_BASE","http://172.18.127.124:40483/v1")
openai_chat_api_base = os.getenv("OPENAI_API_BASE","http://172.18.127.124:40159/v1")
openai_api_key = os.getenv("OPENAI_API_KEY", "gpustack_dee9ca823290886c_5edfc86aeeeceb1e9ee5162941cb2cb5")
embedding_client = OpenAI(
api_key=openai_api_key,
base_url=openai_embedding_api_base,
)
def get_embeddings(embedding_texts: List[str]):
response = embedding_client.embeddings.create(
model="bge-m3",
input=embedding_texts,
)
emebdding_list = response.data[0].embedding
return emebdding_list
conn = connections.connect(host="172.18.127.124", port=19530)
db.using_database("JXTest")
coll_name = 'RAGTest'
search_params = {
"metric_type": 'L2',
"offset": 0,
"ignore_growing": False,
"params": {"nprobe": 16}
}
collection = Collection(coll_name)
collection.load()
query = "手术成功率是多少?"
embeddings = get_embeddings([query])
results = collection.search(
data=[embeddings],
anns_field="vector",
param=search_params,
limit=16,
expr=None,
# output_fields=['m_id', 'embeding', 'desc', 'count'],
output_fields=["text","header_1"],
consistency_level="Strong"
)
collection.release()
# print(results[0].ids)
# print(results[0].distances)
# hit = results[0][0]
# print(hit.entity.get('text'))
print(results)
class Milvus_Database:
def __init__(self,user=os.getenv("Milvus_USER","root"),password=os.getenv("Milvus_PASSWORD","Milvus"),uri=os.getenv("Milvus_URI")):
self.client = MilvusClient(
uri=uri,
token=f"{user}:{password}"
)
def hybrid_search(self,query_vector:List,query,top_k=3,collection_name:str = "MarkdownTest",database_name:str = "JXTest",embedding_rate=0.1)-> List[dict|str]:
"""
"""
try:
self.client.use_database(database_name)
self.load_collection(collection_name)
search_param_1 = {
"data": [query_vector],
"anns_field": "vector",
"param": {
"metric_type": "IP", #COSINE
"params":{"ef":64}#"params": {"nprobe": 10}
},
"limit": top_k
}
request_1 = AnnSearchRequest(**search_param_1)
search_param_2 = {
"data": [query],
"anns_field": "bm_25",
"param": {
"metric_type": "BM25",
},
"limit": top_k
}
request_2 = AnnSearchRequest(**search_param_2)
reqs = [request_1, request_2]
ranker = WeightedRanker(float(embedding_rate), 1-float(embedding_rate))
res = self.client.hybrid_search(
collection_name=collection_name,
reqs=reqs,
ranker=ranker,
limit=top_k,
output_fields=["id","Header_1","Header_2","Header_3","text","resource","img_path","page_idx","origin_id","*"]
#output_fields=["Header_1","Header_2","Header_3","text","resource","img_path","img_information"]
)
search_results = []
search_resources = []
recall_results = []
for i,hits in enumerate(res):
for j,hit in enumerate(hits):
id = hit["entity"].get("id")
search_result_text = hit["entity"].get("text")
distance = hit["distance"]
page_idx = hit["entity"].get("page_idx")
search_result_header_1 = hit["entity"].get("Header_1")
search_result_header_2 = hit["entity"].get("Header_2")
search_result_header_3 = hit["entity"].get("Header_3")
search_result_resource = hit["entity"].get("resource")
search_result_img_path = hit["entity"].get("img_path")
origin_id = hit["entity"].get("origin_id")
header_1_num = len(search_result_header_1) if search_result_header_1 else 0
header_2_num = len(search_result_header_2) if search_result_header_2 else 0
header_3_num = len(search_result_header_3) if search_result_header_3 else 0
#search_result_img_information = hit["entity"].get("img_information")
if search_result_header_1 is None:
search_result_header_2 = ""
else:
search_result_header_1 = f"#{search_result_header_1}"
if search_result_header_2 is None:
search_result_header_2 = ""
else:
search_result_header_2 = f"##{search_result_header_2}"
if search_result_header_3 is None:
search_result_header_3 = ""
else:
search_result_header_3 = f"###{search_result_header_3}"
# search_result = f"#{search_result_header_1}\n{search_result_header_2}\n{search_result_header_3}\n\n{search_result_text}"
search_result = search_result_text
origin_text = search_result[header_1_num+header_2_num+header_3_num+3:]
search_results.append(search_result)
search_resources.append({"resource":search_result_resource,"img_path":search_result_img_path})
#获取片段的下文,并且拼接
next_chunk = self.get_chunk_by_id(collection_name=collection_name,database_name=database_name,id =id)
next_text = next_chunk[0].get("text")
search_result += next_text
origin_text += next_text
# search_resources.append({"resource":search_result_resource,"img_path":search_result_img_path,"img_information":search_result_img_imformation})
recall_results.append({"index":i+j+1,"id":origin_id,"text":search_result,"resource":search_result_resource,"img_path":search_result_img_path,"score":distance,"page_idx":page_idx,"origin_text":origin_text})
#recall_results.append({"index":i+j+1,"text":search_result,"resource":search_result_resource,"img_path":search_result_img_path,"img_information":search_result_img_imformation})
self.release_collection(collection_name)
return search_results, search_resources,recall_results
finally:
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")