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}")