861 lines
37 KiB
Python
861 lines
37 KiB
Python
import json
|
||
import os
|
||
from pymilvus import (
|
||
connections,
|
||
utility,
|
||
FieldSchema,
|
||
CollectionSchema,
|
||
DataType,
|
||
Collection,
|
||
MilvusClient,
|
||
db,
|
||
Function,
|
||
FunctionType,
|
||
AnnSearchRequest,
|
||
WeightedRanker
|
||
)
|
||
from typing import List, Dict, Union, Optional, Any
|
||
from pypinyin import lazy_pinyin
|
||
from typing import List, Dict, Tuple, Optional
|
||
|
||
|
||
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}"
|
||
)
|
||
# self.conn = connections.connect(host="172.18.127.124",port=19530,user=user,password=password)
|
||
|
||
# database = db.create_database(database_name)
|
||
# self.client.list_databases()
|
||
|
||
def create_database(self, database_name) -> List[str]:
|
||
database = self.client.create_database(database_name)
|
||
return self.client.list_databases()
|
||
|
||
def list_databases(self) -> List[str]:
|
||
return self.client.list_databases()
|
||
|
||
def using_database(self, database_name):
|
||
self.client.use_database(database_name)
|
||
|
||
def create_hierarchical_collection(self, collection_name, database_name):
|
||
self.client.use_database(database_name)
|
||
schema = MilvusClient.create_schema(
|
||
auto_id=True,
|
||
enable_dynamic_field=True,
|
||
)
|
||
|
||
analyzer_params_custom = {
|
||
"tokenizer": "jieba",
|
||
"type": "chinese",
|
||
"filter": ["cnalphanumonly"]
|
||
}
|
||
schema.add_field(field_name="id", datatype=DataType.INT64, is_primary=True)
|
||
# 层次召回中的摘要部分
|
||
schema.add_field(field_name="text", datatype=DataType.VARCHAR, max_length=65535, enable_analyzer=True,
|
||
analyzer_params=analyzer_params_custom)
|
||
# 层次召回中的完整文档
|
||
schema.add_field(field_name="origin_text", datatype=DataType.VARCHAR, max_length=65535, enable_analyzer=True,
|
||
analyzer_params=analyzer_params_custom)
|
||
schema.add_field(field_name="bm_25", datatype=DataType.SPARSE_FLOAT_VECTOR)
|
||
# Qwen3-8B-Embedding 维度4096
|
||
schema.add_field(field_name="vector", datatype=DataType.FLOAT_VECTOR, dim=4096, )
|
||
schema.add_field(field_name="resource", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY'),
|
||
schema.add_field(field_name="img_path", datatype=DataType.VARCHAR, max_length=4096, default_value="EMPTY")
|
||
schema.add_field(field_name="page_idx", datatype=DataType.INT64, default_value=0)
|
||
schema.add_field(field_name="origin_id", datatype=DataType.VARCHAR, max_length=256, default_value="EMPTY")
|
||
# schema.add_field(field_name="img_information",datatype=DataType.Json,default_value="EMPTY")
|
||
schema.add_field(field_name="header_1", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
|
||
schema.add_field(field_name="header_2", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
|
||
schema.add_field(field_name="header_3", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
|
||
|
||
bm25_function = Function(
|
||
name="bm25_emb", # Function name
|
||
input_field_names=["text"], # Name of the VARCHAR field containing raw text data
|
||
output_field_names=["bm_25"],
|
||
# Name of the SPARSE_FLOAT_VECTOR field reserved to store generated embeddings
|
||
function_type=FunctionType.BM25,
|
||
)
|
||
schema.add_function(bm25_function)
|
||
index_params = self.client.prepare_index_params()
|
||
# Add indexes
|
||
index_params.add_index(
|
||
field_name="vector",
|
||
index_name="vector_index",
|
||
index_type="HNSW", # "IVF_FLAT". AUTOINDEX
|
||
metric_type="IP", # COSINE
|
||
# params={"nlist": 128},
|
||
)
|
||
|
||
index_params.add_index(
|
||
field_name="bm_25",
|
||
index_name="bm_25_index",
|
||
index_type="SPARSE_INVERTED_INDEX", # Index type for sparse vectors
|
||
metric_type="BM25", # Set to `BM25` when using function to generate sparse vectors
|
||
params={"inverted_index_algo": "DAAT_MAXSCORE"},
|
||
# The ratio of small vector values to be dropped during indexing
|
||
)
|
||
|
||
self.client.create_collection(
|
||
collection_name=collection_name,
|
||
schema=schema,
|
||
index_params=index_params
|
||
)
|
||
self.release_collection(collection_name)
|
||
return "collections created"
|
||
|
||
def create_hybrid_collection(self, collection_name, database_name: str, save_title: bool = True):
|
||
"""
|
||
中文专用混合检索 表
|
||
"""
|
||
self.client.use_database(database_name)
|
||
schema = MilvusClient.create_schema(
|
||
auto_id=True,
|
||
enable_dynamic_field=True,
|
||
)
|
||
|
||
analyzer_params_custom = {
|
||
"tokenizer": "jieba",
|
||
"type": "chinese",
|
||
"filter": ["cnalphanumonly"]
|
||
}
|
||
schema.add_field(field_name="id", datatype=DataType.INT64, is_primary=True)
|
||
schema.add_field(field_name="text", datatype=DataType.VARCHAR, max_length=65535, enable_analyzer=True,
|
||
analyzer_params=analyzer_params_custom)
|
||
schema.add_field(field_name="origin_text", datatype=DataType.VARCHAR, max_length=65535, default_value="EMPTY")
|
||
schema.add_field(field_name="bm_25", datatype=DataType.SPARSE_FLOAT_VECTOR)
|
||
schema.add_field(field_name="vector", datatype=DataType.FLOAT_VECTOR, dim=1024, )
|
||
schema.add_field(field_name="resource", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
|
||
schema.add_field(field_name="img_path", datatype=DataType.VARCHAR, max_length=4096, default_value="EMPTY")
|
||
schema.add_field(field_name="page_idx", datatype=DataType.INT64, default_value=0)
|
||
schema.add_field(field_name="origin_id", datatype=DataType.VARCHAR, max_length=256, default_value="EMPTY")
|
||
# schema.add_field(field_name="img_information",datatype=DataType.Json,default_value="EMPTY")
|
||
if save_title:
|
||
schema.add_field(field_name="header_1", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
|
||
schema.add_field(field_name="header_2", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
|
||
schema.add_field(field_name="header_3", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
|
||
|
||
bm25_function = Function(
|
||
name="bm25_emb", # Function name
|
||
input_field_names=["text"], # Name of the VARCHAR field containing raw text data
|
||
output_field_names=["bm_25"],
|
||
# Name of the SPARSE_FLOAT_VECTOR field reserved to store generated embeddings
|
||
function_type=FunctionType.BM25,
|
||
)
|
||
schema.add_function(bm25_function)
|
||
index_params = self.client.prepare_index_params()
|
||
# Add indexes
|
||
index_params.add_index(
|
||
field_name="vector",
|
||
index_name="vector_index",
|
||
index_type="HNSW", # "IVF_FLAT". AUTOINDEX
|
||
metric_type="IP", # COSINE
|
||
# params={"nlist": 128},
|
||
)
|
||
|
||
index_params.add_index(
|
||
field_name="bm_25",
|
||
index_name="bm_25_index",
|
||
index_type="SPARSE_INVERTED_INDEX", # Index type for sparse vectors
|
||
metric_type="BM25", # Set to `BM25` when using function to generate sparse vectors
|
||
params={"inverted_index_algo": "DAAT_MAXSCORE"},
|
||
# The ratio of small vector values to be dropped during indexing
|
||
)
|
||
|
||
self.client.create_collection(
|
||
collection_name=collection_name,
|
||
schema=schema,
|
||
index_params=index_params
|
||
)
|
||
self.release_collection(collection_name)
|
||
return "collections created"
|
||
|
||
def drop_collection(self, collection_name, database_name: str):
|
||
self.client.use_database(database_name)
|
||
self.load_collection(collection_name=collection_name)
|
||
self.client.drop_collection(collection_name)
|
||
|
||
def load_collection(self, collection_name):
|
||
self.client.load_collection(collection_name=collection_name)
|
||
res = self.client.get_load_state(
|
||
collection_name=collection_name
|
||
)
|
||
return res
|
||
|
||
def release_collection(self, collection_name):
|
||
self.client.release_collection(collection_name=collection_name)
|
||
res = self.client.get_load_state(
|
||
collection_name=collection_name
|
||
)
|
||
return res
|
||
|
||
def delete_resource(self, collection_name, resource, database_name):
|
||
self.client.use_database(database_name)
|
||
self.load_collection(collection_name)
|
||
res = self.client.delete(
|
||
collection_name=collection_name,
|
||
filter="resource == '{}'".format(resource),
|
||
)
|
||
return res
|
||
|
||
def has_collection(self, collection_name, database_name):
|
||
"""
|
||
Before using list collection, You need to using using_database method first
|
||
"""
|
||
self.using_database(database_name=database_name)
|
||
return self.client.has_collection(collection_name)
|
||
|
||
def list_collections(self, database_name):
|
||
"""
|
||
Before using list collection, You need to using using_database method first
|
||
"""
|
||
self.client.use_database(database_name)
|
||
return self.client.list_collections()
|
||
|
||
def hierarchical_insert(self, params, collection_name: str, database_name: str):
|
||
try:
|
||
self.client.use_database(database_name)
|
||
self.load_collection(collection_name)
|
||
res = self.client.insert(
|
||
collection_name=collection_name,
|
||
data=params
|
||
)
|
||
return res
|
||
finally:
|
||
try:
|
||
self.client.release_collection(collection_name)
|
||
except Exception as e:
|
||
print(f"释放集合失败: {e}")
|
||
|
||
def hybrid_insert(self, params, collection_name: str, database_name: str):
|
||
try:
|
||
self.client.use_database(database_name)
|
||
self.load_collection(collection_name)
|
||
res = self.client.insert(
|
||
collection_name=collection_name,
|
||
data=params
|
||
)
|
||
return res
|
||
finally:
|
||
try:
|
||
self.client.release_collection(collection_name)
|
||
except Exception as e:
|
||
print(f"释放集合失败: {e}")
|
||
|
||
# def insert(self, params,collection_name:str = "MarkdownTest",database_name:str = "JXTest"):
|
||
# """
|
||
# params: List[List]
|
||
# 例子: [[id, header_1, header_2, header_3, text, vector]]
|
||
# """
|
||
# # self.client.insert(collection_name=collection_name, data=params)
|
||
# db.using_database(database_name)
|
||
# collection = Collection(collection_name)
|
||
# mr = collection.insert(params)
|
||
# return mr
|
||
def get_all_chunk_by_resource(self, collection_name, database_name, resource):
|
||
"""
|
||
return List :
|
||
"""
|
||
try:
|
||
self.client.using_database(db_name=database_name)
|
||
if not self.client.has_collection(collection_name):
|
||
return []
|
||
else:
|
||
self.client.load_collection(collection_name)
|
||
iterator = self.client.query_iterator(
|
||
collection_name=collection_name,
|
||
batch_size=5,
|
||
filter=f"resource like \"{resource}\"",
|
||
output_fields=["id", "text", "resource", "Header_1", "Header_2", "Header_3"],
|
||
)
|
||
|
||
results = []
|
||
|
||
while True:
|
||
result = iterator.next()
|
||
if not result:
|
||
iterator.close()
|
||
break
|
||
results += result
|
||
return results
|
||
finally:
|
||
# 3. 无论成功还是失败,最终释放集合
|
||
try:
|
||
self.client.release_collection(collection_name)
|
||
except Exception as e:
|
||
print(f"释放集合失败: {e}")
|
||
|
||
def get_chunk_by_resource(self, collection_name, database_name, resource):
|
||
"""
|
||
return List :
|
||
"""
|
||
try:
|
||
self.client.using_database(db_name=database_name)
|
||
if not self.client.has_collection(collection_name):
|
||
return []
|
||
else:
|
||
self.client.load_collection(collection_name)
|
||
iterator = self.client.query_iterator(
|
||
collection_name=collection_name,
|
||
batch_size=5,
|
||
filter=f"resource like \"{resource}\"",
|
||
output_fields=["id", "text", "resource"],
|
||
)
|
||
|
||
results = []
|
||
|
||
while True:
|
||
result = iterator.next()
|
||
if not result:
|
||
iterator.close()
|
||
break
|
||
results += result
|
||
return results
|
||
finally:
|
||
# 3. 无论成功还是失败,最终释放集合
|
||
try:
|
||
self.client.release_collection(collection_name)
|
||
except Exception as e:
|
||
print(f"释放集合失败: {e}")
|
||
|
||
def get_chunk_by_id(self, collection_name, database_name, id):
|
||
"""
|
||
return List :
|
||
"""
|
||
try:
|
||
self.client.using_database(db_name=database_name)
|
||
if not self.client.has_collection(collection_name):
|
||
return []
|
||
else:
|
||
self.client.load_collection(collection_name)
|
||
iterator = self.client.query_iterator(
|
||
collection_name=collection_name,
|
||
batch_size=5,
|
||
filter=f"id=={id + 1}",
|
||
output_fields=["id", "text", "resource"],
|
||
)
|
||
|
||
results = []
|
||
|
||
while True:
|
||
result = iterator.next()
|
||
if not result:
|
||
iterator.close()
|
||
break
|
||
results += result
|
||
return results
|
||
finally:
|
||
# 3. 无论成功还是失败,最终释放集合
|
||
try:
|
||
self.client.release_collection(collection_name)
|
||
except Exception as e:
|
||
print(f"释放集合失败: {e}")
|
||
|
||
def get_text_by_resource(self, collection_name, database_name, resource, logger):
|
||
"""
|
||
return List :
|
||
"""
|
||
|
||
# def prepare_milvus_filter( resource):
|
||
# """安全构造 Milvus 过滤条件"""
|
||
# escaped_resource = json.dumps(resource)[1:-1]
|
||
# return f'resource="{escaped_resource}"'
|
||
def prepare_milvus_filter(resource):
|
||
"""安全构造 Milvus 过滤条件"""
|
||
# 替换资源名称中的双引号
|
||
safe_resource = resource.replace('"', '\\"')
|
||
return f'resource == "{safe_resource}"'
|
||
|
||
self.client.using_database(db_name=database_name)
|
||
|
||
# 检查集合是否存在
|
||
if not self.client.has_collection(collection_name):
|
||
logger.info(f"集合 {collection_name} 不存在")
|
||
return []
|
||
else:
|
||
self.client.load_collection(collection_name)
|
||
try:
|
||
# 2. 执行查询
|
||
filter_expr = prepare_milvus_filter(resource)
|
||
iterator = self.client.query_iterator(
|
||
collection_name=collection_name,
|
||
batch_size=5,
|
||
filter=filter_expr,
|
||
output_fields=["id", "resource", "text"],
|
||
)
|
||
|
||
results = []
|
||
while True:
|
||
try:
|
||
result = iterator.next()
|
||
if not result:
|
||
break
|
||
results.extend(result)
|
||
except StopIteration:
|
||
break
|
||
except Exception as e:
|
||
print(f"查询迭代出错: {e}")
|
||
break
|
||
|
||
return results
|
||
|
||
except Exception as e:
|
||
print(f"查询失败: {e}")
|
||
return []
|
||
|
||
finally:
|
||
# 3. 无论成功还是失败,最终释放集合
|
||
try:
|
||
self.client.release_collection(collection_name)
|
||
except Exception as e:
|
||
print(f"释放集合失败: {e}")
|
||
|
||
def embedding_search(self, query_vector: List, top_k=3, collection_name: str = "MarkdownTest",
|
||
database_name: str = "JXTest") -> List[dict | str]:
|
||
"""
|
||
纯向量检索
|
||
"""
|
||
self.client.use_database(database_name)
|
||
self.load_collection(collection_name)
|
||
res = self.client.search(
|
||
collection_name=collection_name,
|
||
anns_field="vector",
|
||
limit=top_k,
|
||
data=[query_vector],
|
||
search_params={
|
||
"metric_type": "IP", # COSINE
|
||
"params": {"ef": 64} # "params": {"nprobe": 10}
|
||
},
|
||
output_fields=["id", "Header_1", "Header_2", "Header_3", "text", "resource", "img_path", "page_idx",
|
||
"origin_id", "*"])
|
||
search_results = []
|
||
search_resources = []
|
||
recall_results = []
|
||
try:
|
||
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")
|
||
dynamic_field_name = hit["entity"].get_properties()
|
||
# search_score = hit["entity"].get("score")
|
||
# 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
|
||
search_results.append(search_result)
|
||
search_resources.append({"resource": search_result_resource, "img_path": search_result_img_path})
|
||
# search_resources.append({"resource":search_result_resource,"img_path":search_result_img_path,"img_information":search_result_img_imformation})
|
||
recall_results.append(
|
||
{"id": origin_id, "index": i + j + 1, "text": search_result, "resource": search_result_resource,
|
||
"img_path": search_result_img_path, "score": distance, "page_idx": page_idx,
|
||
"dynamic_field_name": dynamic_field_name})
|
||
# 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})
|
||
return search_results, search_resources, recall_results
|
||
finally:
|
||
# 3. 无论成功还是失败,最终释放集合
|
||
try:
|
||
self.client.release_collection(collection_name)
|
||
except Exception as e:
|
||
print(f"释放集合失败: {e}")
|
||
|
||
def hierarchical_search(self, query_vector: List, query, top_k=3, collection_name: str = "MarkdownTest",
|
||
database_name: str = "JXTest") -> 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(0.3, 0.7)
|
||
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")
|
||
text = hit["entity"].get("text")
|
||
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")
|
||
origin_text = hit["entity".get("origin_text")]
|
||
# 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
|
||
search_results.append(search_result)
|
||
search_resources.append({"resource": search_result_resource, "img_path": search_result_img_path})
|
||
# 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":origin_text,"resource":search_result_resource,"img_path":search_result_img_path,"score":distance,"page_idx":page_idx})
|
||
recall_results.append(
|
||
{"index": i + j + 1, "id": origin_id, "text": text, "resource": search_result_resource,
|
||
"img_path": search_result_img_path, "score": distance, "page_idx": page_idx})
|
||
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}")
|
||
|
||
def hybrid_search(self, query_vector: List, query, top_k=3, collection_name: str = "MarkdownTest",
|
||
database_name: str = "JXTest", embedding_rate=0.3) -> 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", "origin_text"]
|
||
# 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
|
||
origin_text = hit["entity"].get("origin_text")
|
||
# 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}")
|
||
|
||
# return search_results
|
||
|
||
def hybrid_search_batch(
|
||
self,
|
||
query_vectors: List[List[float]],
|
||
queries: List[str],
|
||
top_k: int = 3,
|
||
collection_name: str = "MarkdownTest",
|
||
database_name: str = "JXTest",
|
||
embedding_rate: float = 0.3
|
||
) -> Tuple[List[List[Dict]], List[List[Dict]], List[List[Dict]]]:
|
||
"""
|
||
执行批量混合检索(稠密向量 + BM25 稀疏检索),对每个查询分别进行 hybrid search。
|
||
|
||
Args:
|
||
query_vectors (List[List[float]]): 每个查询对应的稠密向量列表,形状为 [N, dim]
|
||
queries (List[str]): 原始文本查询列表,长度为 N
|
||
top_k (int): 返回每个查询的 top-k 结果
|
||
collection_name (str): 集合名称
|
||
database_name (str): 数据库名称
|
||
embedding_rate (float): 稠密向量检索结果的权重(BM25 权重 = 1 - embedding_rate)
|
||
|
||
Returns:
|
||
Tuple[
|
||
List[List[str]]: 每个查询对应的结果文本列表(外层 list 是 query 数量,内层是 top-k 结果),
|
||
List[List[dict]]: 每个查询对应的资源信息列表 ,
|
||
List[List[dict]]: 每个查询对应的完整召回详情列表(含 score、id、origin_text 等)
|
||
]
|
||
"""
|
||
if len(query_vectors) != len(queries):
|
||
raise ValueError("query_vectors 和 queries 的数量必须一致")
|
||
|
||
try:
|
||
self.client.use_database(database_name)
|
||
self.load_collection(collection_name)
|
||
|
||
all_search_results = [] # [[res1_q1, res2_q1], [res1_q2, ...]]
|
||
all_search_resources = []
|
||
all_recall_results = []
|
||
|
||
for i, (vec, text) in enumerate(zip(query_vectors, queries)):
|
||
try:
|
||
# 构建稠密向量请求
|
||
req1 = AnnSearchRequest(
|
||
data=[vec],
|
||
anns_field="vector",
|
||
param={"metric_type": "IP", "params": {"ef": 64}},
|
||
limit=top_k
|
||
)
|
||
|
||
# 构建稀疏向量请求(BM25)
|
||
req2 = AnnSearchRequest(
|
||
data=[text],
|
||
anns_field="bm_25",
|
||
param={"metric_type": "BM25"},
|
||
limit=top_k
|
||
)
|
||
|
||
# 混合搜索
|
||
ranker = WeightedRanker(embedding_rate, 1 - embedding_rate)
|
||
res = self.client.hybrid_search(
|
||
collection_name=collection_name,
|
||
reqs=[req1, req2],
|
||
ranker=ranker,
|
||
limit=top_k,
|
||
output_fields=[
|
||
"id", "Header_1", "Header_2", "Header_3", "text", "resource",
|
||
"img_path", "page_idx", "origin_id", "origin_text"
|
||
],
|
||
timeout=30
|
||
)
|
||
|
||
# 解析单个查询的结果
|
||
search_results = []
|
||
search_resources = []
|
||
recall_results = []
|
||
|
||
for j, hit in enumerate(res[0]): # res 是 List[List[Hit]], res[0] 对应第一个 request 的 hits
|
||
entity = hit["entity"]
|
||
id_val = entity.get("id")
|
||
origin_id = entity.get("origin_id")
|
||
text_val = entity.get("text") or ""
|
||
header_1 = entity.get("Header_1") or ""
|
||
header_2 = entity.get("Header_2") or ""
|
||
header_3 = entity.get("Header_3") or ""
|
||
resource = entity.get("resource")
|
||
img_path = entity.get("img_path")
|
||
page_idx = entity.get("page_idx")
|
||
distance = hit["distance"]
|
||
origin_text = entity.get("origin_text") or text_val
|
||
|
||
# 格式化标题
|
||
header_1 = f"#{header_1}" if header_1 else ""
|
||
header_2 = f"##{header_2}" if header_2 else ""
|
||
header_3 = f"###{header_3}" if header_3 else ""
|
||
|
||
# 组合内容(可根据需要调整)
|
||
combined_text = "\n".join(filter(None, [header_1, header_2, header_3, text_val])).strip()
|
||
|
||
search_results.append(combined_text)
|
||
search_resources.append({"resource": resource, "img_path": img_path})
|
||
|
||
recall_results.append({
|
||
"keyword": text,
|
||
"index": j + 1,
|
||
"id": origin_id,
|
||
"text": combined_text,
|
||
"resource": resource,
|
||
"img_path": img_path,
|
||
"score": distance,
|
||
"page_idx": page_idx,
|
||
"origin_text": origin_text
|
||
})
|
||
|
||
all_search_results.append(search_results)
|
||
all_search_resources.append(search_resources)
|
||
all_recall_results.append(recall_results)
|
||
|
||
except Exception as e:
|
||
print(f"第 {i + 1} 个查询 '{text}' 检索失败: {e}")
|
||
# 失败时返回空列表占位
|
||
all_search_results.append([])
|
||
all_search_resources.append([])
|
||
all_recall_results.append([])
|
||
|
||
return all_search_results, all_search_resources, all_recall_results
|
||
|
||
except Exception as e:
|
||
print(f"批量混合检索整体失败: {e}")
|
||
raise
|
||
|
||
finally:
|
||
try:
|
||
self.client.release_collection(collection_name)
|
||
except Exception as e:
|
||
print(f"释放集合失败: {e}")
|
||
# return search_results
|
||
|
||
def close(self):
|
||
self.client.close()
|
||
|
||
|
||
def sanitize_collection_name(name: str) -> str:
|
||
"""
|
||
中文转换为拼音
|
||
"""
|
||
|
||
def is_valid_milvus_collection_name(name: str) -> bool:
|
||
"""检查名称是否符合 Milvus 集合命名规则"""
|
||
import re
|
||
pattern = r'^[a-zA-Z_][a-zA-Z0-9_]*$' # 首字符必须是字母或_,后续可以是字母、数字、_
|
||
return bool(re.fullmatch(pattern, name))
|
||
|
||
if is_valid_milvus_collection_name(name):
|
||
return name
|
||
pinyin = "_".join(lazy_pinyin(name))
|
||
# 确保首字符是字母或下划线
|
||
if not pinyin[0].isalpha() and pinyin[0] != "_":
|
||
pinyin = f"_{pinyin}"
|
||
return pinyin
|
||
|
||
|
||
##应用示例
|
||
# if __name__ == '__main__':
|
||
# #创建collection
|
||
# database = Milvus_Database("XIAN")
|
||
|
||
# result = database.create_connection("MarkdownTest",drop_=False)
|
||
# print(result)
|
||
# print(database.list_partitions("MarkdownTest"))
|
||
|
||
|
||
# 创建db
|
||
# from pymilvus import connections, db
|
||
# #
|
||
# conn = Milvus_Database(user='root',
|
||
# password='Milvus',
|
||
# uri="http://192.168.0.46:19530")
|
||
# print(conn.list_collections("XIAN"))
|
||
|
||
# conn = connections.connect(user='root',
|
||
# password='Milvus',
|
||
# uri="http://172.18.30.165:19530")
|
||
|
||
# database = db.create_database("JXTest")
|
||
# db.using_database("JXTest")
|
||
|