tsystems/visual_document_retrieval
2
1from qdrant_client import QdrantClient2from qdrant_client.http import models3from tqdm import tqdm4import os5import time6import numpy as np7from loguru import logger8import stamina9from typing import Any, List, Tuple, Type, Literal, Optional, Union, Dict10 11class MyQdrantClient:12 def __init__(self, path: str):13 self.qdrant_client = QdrantClient(path=path)14 logger.debug(f"Qdrant client created at {path}")15 16 def create_collection(self, collection_name: str, vector_dim: int = 128, vector_type: str = "colbert"):17 if vector_type == "colbert":18 self.qdrant_client.create_collection(19 collection_name=collection_name,20 on_disk_payload=True, # store the payload on disk21 vectors_config=models.VectorParams(22 size=vector_dim,23 distance=models.Distance.COSINE,24 on_disk=True, # move original vectors to disk25 multivector_config=models.MultiVectorConfig(26 comparator=models.MultiVectorComparator.MAX_SIM27 ),28 #quantization_config=models.BinaryQuantization(29 #binary=models.BinaryQuantizationConfig(30 # always_ram=True # keep only quantized vectors in RAM31 # ),32 #),33 ),34 )35 elif vector_type == "dense":36 self.qdrant_client.create_collection(37 collection_name=collection_name,38 on_disk_payload=True, # store the payload on disk39 vectors_config=models.VectorParams(40 size=vector_dim,41 distance=models.Distance.COSINE,42 on_disk=True, # move original vectors to disk43 ),44 )45 else:46 raise ValueError(f"Vector type {vector_type} not supported")47 48 logger.debug(f"Qdrant collection of type {vector_type} : {collection_name} created")49 50 def delete_collection(self, collection_name: str):51 self.qdrant_client.delete_collection(collection_name=collection_name)52 53 @stamina.retry(on=Exception, attempts=3) # retry mechanism if an exception occurs during the operation54 def upsert_to_qdrant(self, batch, collection_name: str):55 try:56 self.qdrant_client.upsert(57 collection_name=collection_name,58 points=batch,59 wait=False,60 )61 except Exception as e:62 logger.error(f"Error during upsert: {e}")63 return False64 return True65 66 def upsert_multivector(self, index: int, multivector_input_list: list[Any], collection_name: str):67 try:68 points = []69 for j, multivector in enumerate(multivector_input_list):70 points.append(71 models.PointStruct(72 id=index + j, # we just use the index as the ID73 vector=multivector, # This is now a list of vectors74 payload={75 "source": "user uploaded data"76 }, # can also add other metadata/data77 )78 )79 # Upload points to Qdrant80 81 self.upsert_to_qdrant(points, collection_name)82 except Exception as e:83 logger.error(f"Vector DB client - error during upsert: {e}")84 85 def query_multivector(self, multivector_input, collection_name: str, top_k:int=10) -> list[int]:86 try:87 #logger.debug(f"Number of vector: {len(multivector_input)}")88 #logger.debug(f"Vector dim: {len(multivector_input[0])}")89 90 start_time = time.time()91 search_result = self.qdrant_client.query_points(92 collection_name=collection_name,93 query=multivector_input,94 limit=top_k,95 # timeout=100,96 # search_params=models.SearchParams(97 # quantization=models.QuantizationSearchParams(98 # ignore=False,99 # rescore=True,100 # oversampling=2.0,101 # )102 # )103 )104 end_time = time.time()105 elapsed_time = end_time - start_time106 logger.debug(f"Search completed in {elapsed_time:.4f} seconds")107 108 result = [x.id for x in search_result.points]109 return result110 111 except Exception as e:112 logger.error(f"Error during query: {e}")113 return None114 115 def __del__(self):116 self.qdrant_client.close()117 