158 lines
6.7 KiB
Python
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}") |