CoolFace
Apppublic

tsystems/visual_document_retrieval

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
2likes
qdrant_db.py117 linesDownload Raw Back to app
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