2024-07-09 05:04:20 +08:00
|
|
|
import os
|
|
|
|
|
|
2024-07-28 16:16:52 +08:00
|
|
|
from src.models.embedding import EmbeddingModel
|
2024-07-09 05:04:20 +08:00
|
|
|
from pymilvus import MilvusClient
|
2024-07-28 16:16:52 +08:00
|
|
|
from src.utils import setup_logger, hashstr
|
2024-07-14 23:59:52 +08:00
|
|
|
logger = setup_logger("KnowledgeBase")
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
|
|
|
|
2024-07-14 23:59:52 +08:00
|
|
|
class KnowledgeBase:
|
2024-07-09 05:04:20 +08:00
|
|
|
|
2024-07-28 16:16:52 +08:00
|
|
|
def __init__(self, config=None, embed_model=None) -> None:
|
2024-07-09 05:04:20 +08:00
|
|
|
self.config = config
|
|
|
|
|
self._init_config(config)
|
2024-07-28 16:16:52 +08:00
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
assert embed_model, "embed_model=None"
|
|
|
|
|
self.embed_model = embed_model
|
2024-08-07 17:24:46 +08:00
|
|
|
|
2024-07-28 16:16:52 +08:00
|
|
|
self.client = MilvusClient(self.milvus_path)
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
|
|
|
def _init_config(self, config):
|
|
|
|
|
self.vector_dim = 1024 # 暂时不知道这个和 embedding model 的 embedding 大小有什么关系
|
2024-07-28 16:16:52 +08:00
|
|
|
self.milvus_path = os.path.join(config.save_dir, "data/vector_base/milvus.db")
|
|
|
|
|
os.makedirs(os.path.dirname(self.milvus_path), exist_ok=True)
|
2024-07-09 05:04:20 +08:00
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
def get_collection_names(self):
|
|
|
|
|
return self.client.list_collections()
|
|
|
|
|
|
2024-07-14 23:59:52 +08:00
|
|
|
def get_collections(self):
|
|
|
|
|
collections_name = self.client.list_collections()
|
|
|
|
|
collections = []
|
|
|
|
|
for collection_name in collections_name:
|
|
|
|
|
collection = self.get_collection_info(collection_name)
|
|
|
|
|
collections.append(collection)
|
|
|
|
|
|
|
|
|
|
return collections
|
|
|
|
|
|
|
|
|
|
def get_collection_info(self, collection_name):
|
|
|
|
|
collection = self.client.describe_collection(collection_name)
|
|
|
|
|
collection.update(self.client.get_collection_stats(collection_name))
|
|
|
|
|
# collection["id"] = hashstr(collection_name)
|
|
|
|
|
return collection
|
|
|
|
|
|
|
|
|
|
def add_collection(self, collection_name):
|
|
|
|
|
if self.client.has_collection(collection_name=collection_name):
|
|
|
|
|
logger.warning(f"Collection {collection_name} already exists, drop it")
|
|
|
|
|
self.client.drop_collection(collection_name=collection_name)
|
|
|
|
|
|
|
|
|
|
self.client.create_collection(
|
|
|
|
|
collection_name=collection_name,
|
|
|
|
|
dimension=self.vector_dim, # The vectors we will use in this demo has 768 dimensions
|
|
|
|
|
)
|
|
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
def add_documents(self, docs, collection_name, **kwargs):
|
|
|
|
|
"""添加已经分块之后的文本"""
|
2024-07-14 23:59:52 +08:00
|
|
|
# 检查 collection 是否存在
|
2024-07-16 18:14:27 +08:00
|
|
|
import random
|
2024-07-14 23:59:52 +08:00
|
|
|
if not self.client.has_collection(collection_name=collection_name):
|
|
|
|
|
logger.warning(f"Collection {collection_name} not found, create it")
|
|
|
|
|
self.add_collection(collection_name)
|
|
|
|
|
|
2024-07-09 05:04:20 +08:00
|
|
|
vectors = self.embed_model.encode(docs)
|
|
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
data = [{
|
|
|
|
|
"id": int(random.random() * 1e12),
|
|
|
|
|
"vector": vectors[i],
|
|
|
|
|
"text": docs[i],
|
|
|
|
|
"hash": hashstr(docs[i] + str(random.random())),
|
|
|
|
|
**kwargs} for i in range(len(vectors))]
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
|
|
|
res = self.client.insert(collection_name=collection_name, data=data)
|
|
|
|
|
return res
|
|
|
|
|
|
2024-07-17 18:52:20 +08:00
|
|
|
def search(self, query, collection_name, limit=3):
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
|
|
|
query_vectors = self.embed_model.encode_queries([query])
|
|
|
|
|
|
|
|
|
|
res = self.client.search(
|
|
|
|
|
collection_name=collection_name, # target collection
|
|
|
|
|
data=query_vectors, # query vectors
|
|
|
|
|
limit=limit, # number of returned entities
|
2024-07-28 16:16:52 +08:00
|
|
|
output_fields=["text", "file_id"], # specifies fields to be returned
|
2024-07-09 05:04:20 +08:00
|
|
|
)
|
|
|
|
|
|
2024-07-10 12:56:53 +08:00
|
|
|
return res[0] # 因为 query 只有一个
|
2024-07-09 05:04:20 +08:00
|
|
|
|
2024-07-14 23:59:52 +08:00
|
|
|
def examples(self, collection_name, limit=20):
|
|
|
|
|
res = self.client.query(
|
|
|
|
|
collection_name=collection_name,
|
|
|
|
|
limit=10,
|
|
|
|
|
output_fields=["id", "text"],
|
|
|
|
|
)
|
|
|
|
|
return res
|
|
|
|
|
|
|
|
|
|
def search_by_id(self, collection_name, id, output_fields=["id", "text"]):
|
|
|
|
|
res = self.client.get(collection_name, id, output_fields=output_fields)
|
|
|
|
|
return res
|