Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
baiduvectordb.py439 linesDownload Raw Back to vectorstores
1"""Wrapper around the Baidu vector database."""2 3from __future__ import annotations4 5import json6import logging7import time8from typing import Any, Dict, Iterable, List, Optional, Tuple9 10import numpy as np11from langchain_core.documents import Document12from langchain_core.embeddings import Embeddings13from langchain_core.utils import guard_import14from langchain_core.vectorstores import VectorStore15 16from langchain_community.vectorstores.utils import maximal_marginal_relevance17 18logger = logging.getLogger(__name__)19 20 21class ConnectionParams:22    """Baidu VectorDB Connection params.23 24    See the following documentation for details:25    https://cloud.baidu.com/doc/VDB/s/6lrsob0wy26 27    Attribute:28        endpoint (str) : The access address of the vector database server29            that the client needs to connect to.30        api_key (str): API key for client to access the vector database server,31            which is used for authentication.32        account (str) : Account for client to access the vector database server.33        connection_timeout_in_mills (int) : Request Timeout.34    """35 36    def __init__(37        self,38        endpoint: str,39        api_key: str,40        account: str = "root",41        connection_timeout_in_mills: int = 50 * 1000,42    ):43        self.endpoint = endpoint44        self.api_key = api_key45        self.account = account46        self.connection_timeout_in_mills = connection_timeout_in_mills47 48 49class TableParams:50    """Baidu VectorDB table params.51 52    See the following documentation for details:53    https://cloud.baidu.com/doc/VDB/s/mlrsob0p654    """55 56    def __init__(57        self,58        dimension: int,59        replication: int = 3,60        partition: int = 1,61        index_type: str = "HNSW",62        metric_type: str = "L2",63        params: Optional[Dict] = None,64    ):65        self.dimension = dimension66        self.replication = replication67        self.partition = partition68        self.index_type = index_type69        self.metric_type = metric_type70        self.params = params71 72 73class BaiduVectorDB(VectorStore):74    """Baidu VectorDB as a vector store.75 76    In order to use this you need to have a database instance.77    See the following documentation for details:78    https://cloud.baidu.com/doc/VDB/index.html79    """80 81    field_id: str = "id"82    field_vector: str = "vector"83    field_text: str = "text"84    field_metadata: str = "metadata"85 86    index_vector: str = "vector_idx"87 88    def __init__(89        self,90        embedding: Embeddings,91        connection_params: ConnectionParams,92        table_params: TableParams = TableParams(128),93        database_name: str = "LangChainDatabase",94        table_name: str = "LangChainTable",95        drop_old: Optional[bool] = False,96    ):97        pymochow = guard_import("pymochow")98        configuration = guard_import("pymochow.configuration")99        auth = guard_import("pymochow.auth.bce_credentials")100        self.mochowtable = guard_import("pymochow.model.table")101        self.mochowenum = guard_import("pymochow.model.enum")102        self.embedding_func = embedding103        self.table_params = table_params104        config = configuration.Configuration(105            credentials=auth.BceCredentials(106                connection_params.account, connection_params.api_key107            ),108            endpoint=connection_params.endpoint,109            connection_timeout_in_mills=connection_params.connection_timeout_in_mills,110        )111        self.vdb_client = pymochow.MochowClient(config)112        db_list = self.vdb_client.list_databases()113        db_exist: bool = False114        for db in db_list:115            if database_name == db.database_name:116                db_exist = True117                break118        if db_exist:119            self.database = self.vdb_client.database(database_name)120        else:121            self.database = self.vdb_client.create_database(database_name)122        try:123            self.table = self.database.describe_table(table_name)124            if drop_old:125                self.database.drop_table(table_name)126                self._create_table(table_name)127        except pymochow.exception.ServerError:128            self._create_table(table_name)129 130    def _create_table(self, table_name: str) -> None:131        schema = guard_import("pymochow.model.schema")132        index_type = None133        for k, v in self.mochowenum.IndexType.__members__.items():134            if k == self.table_params.index_type:135                index_type = v136        if index_type is None:137            raise ValueError("unsupported index_type")138        metric_type = None139        for k, v in self.mochowenum.MetricType.__members__.items():140            if k == self.table_params.metric_type:141                metric_type = v142        if metric_type is None:143            raise ValueError("unsupported metric_type")144        if self.table_params.params is None:145            params = schema.HNSWParams(m=16, efconstruction=200)146        else:147            params = schema.HNSWParams(148                m=self.table_params.params.get("M", 16),149                efconstruction=self.table_params.params.get("efConstruction", 200),150            )151        fields = []152        fields.append(153            schema.Field(154                self.field_id,155                self.mochowenum.FieldType.STRING,156                primary_key=True,157                partition_key=True,158                auto_increment=False,159                not_null=True,160            )161        )162        fields.append(163            schema.Field(164                self.field_vector,165                self.mochowenum.FieldType.FLOAT_VECTOR,166                dimension=self.table_params.dimension,167                not_null=True,168            )169        )170        fields.append(schema.Field(self.field_text, self.mochowenum.FieldType.STRING))171        fields.append(172            schema.Field(self.field_metadata, self.mochowenum.FieldType.STRING)173        )174        indexes = []175        indexes.append(176            schema.VectorIndex(177                index_name=self.index_vector,178                index_type=index_type,179                field=self.field_vector,180                metric_type=metric_type,181                params=params,182            )183        )184 185        self.table = self.database.create_table(186            table_name=table_name,187            replication=self.table_params.replication,188            partition=self.mochowtable.Partition(189                partition_num=self.table_params.partition190            ),191            schema=schema.Schema(fields=fields, indexes=indexes),192        )193 194        while True:195            time.sleep(1)196            table = self.database.describe_table(table_name)197            if table.state == self.mochowenum.TableState.NORMAL:198                break199 200    @property201    def embeddings(self) -> Embeddings:202        return self.embedding_func203 204    @classmethod205    def from_texts(206        cls,207        texts: List[str],208        embedding: Embeddings,209        metadatas: Optional[List[dict]] = None,210        connection_params: Optional[ConnectionParams] = None,211        table_params: Optional[TableParams] = None,212        database_name: str = "LangChainDatabase",213        table_name: str = "LangChainTable",214        drop_old: Optional[bool] = False,215        **kwargs: Any,216    ) -> BaiduVectorDB:217        """Create a table, indexes it with HNSW, and insert data."""218        if len(texts) == 0:219            raise ValueError("texts is empty")220        if connection_params is None:221            raise ValueError("connection_params is empty")222        try:223            embeddings = embedding.embed_documents(texts[0:1])224        except NotImplementedError:225            embeddings = [embedding.embed_query(texts[0])]226        dimension = len(embeddings[0])227        if table_params is None:228            table_params = TableParams(dimension=dimension)229        else:230            table_params.dimension = dimension231        vector_db = cls(232            embedding=embedding,233            connection_params=connection_params,234            table_params=table_params,235            database_name=database_name,236            table_name=table_name,237            drop_old=drop_old,238        )239        vector_db.add_texts(texts=texts, metadatas=metadatas)240        return vector_db241 242    def add_texts(243        self,244        texts: Iterable[str],245        metadatas: Optional[List[dict]] = None,246        batch_size: int = 1000,247        **kwargs: Any,248    ) -> List[str]:249        """Insert text data into Baidu VectorDB."""250        texts = list(texts)251        try:252            embeddings = self.embedding_func.embed_documents(texts)253        except NotImplementedError:254            embeddings = [self.embedding_func.embed_query(x) for x in texts]255        if len(embeddings) == 0:256            logger.debug("Nothing to insert, skipping.")257            return []258        pks: list[str] = []259        total_count = len(embeddings)260        for start in range(0, total_count, batch_size):261            # Grab end index262            rows = []263            end = min(start + batch_size, total_count)264            for id in range(start, end, 1):265                metadata = "{}"266                if metadatas is not None:267                    metadata = json.dumps(metadatas[id])268                row = self.mochowtable.Row(269                    id="{}-{}-{}".format(time.time_ns(), hash(texts[id]), id),270                    vector=[float(num) for num in embeddings[id]],271                    text=texts[id],272                    metadata=metadata,273                )274                rows.append(row)275                pks.append(str(id))276            self.table.upsert(rows=rows)277        # need rebuild vindex after upsert278        self.table.rebuild_index(self.index_vector)279        while True:280            time.sleep(2)281            index = self.table.describe_index(self.index_vector)282            if index.state == self.mochowenum.IndexState.NORMAL:283                break284        return pks285 286    def similarity_search(287        self,288        query: str,289        k: int = 4,290        param: Optional[dict] = None,291        expr: Optional[str] = None,292        **kwargs: Any,293    ) -> List[Document]:294        """Perform a similarity search against the query string."""295        res = self.similarity_search_with_score(296            query=query, k=k, param=param, expr=expr, **kwargs297        )298        return [doc for doc, _ in res]299 300    def similarity_search_with_score(301        self,302        query: str,303        k: int = 4,304        param: Optional[dict] = None,305        expr: Optional[str] = None,306        **kwargs: Any,307    ) -> List[Tuple[Document, float]]:308        """Perform a search on a query string and return results with score."""309        # Embed the query text.310        embedding = self.embedding_func.embed_query(query)311        res = self._similarity_search_with_score(312            embedding=embedding, k=k, param=param, expr=expr, **kwargs313        )314        return res315 316    def similarity_search_by_vector(317        self,318        embedding: List[float],319        k: int = 4,320        param: Optional[dict] = None,321        expr: Optional[str] = None,322        **kwargs: Any,323    ) -> List[Document]:324        """Perform a similarity search against the query string."""325        res = self._similarity_search_with_score(326            embedding=embedding, k=k, param=param, expr=expr, **kwargs327        )328        return [doc for doc, _ in res]329 330    def _similarity_search_with_score(331        self,332        embedding: List[float],333        k: int = 4,334        param: Optional[dict] = None,335        expr: Optional[str] = None,336        **kwargs: Any,337    ) -> List[Tuple[Document, float]]:338        """Perform a search on a query string and return results with score."""339        ef = 10 if param is None else param.get("ef", 10)340 341        anns = self.mochowtable.AnnSearch(342            vector_field=self.field_vector,343            vector_floats=[float(num) for num in embedding],344            params=self.mochowtable.HNSWSearchParams(ef=ef, limit=k),345            filter=expr,346        )347        res = self.table.search(anns=anns)348 349        rows = [[item] for item in res.rows]350        # Organize results.351        ret: List[Tuple[Document, float]] = []352        if rows is None or len(rows) == 0:353            return ret354        for row in rows:355            for result in row:356                row_data = result.get("row", {})357                meta = row_data.get(self.field_metadata)358                if meta is not None:359                    meta = json.loads(meta)360                doc = Document(361                    page_content=row_data.get(self.field_text), metadata=meta362                )363                pair = (doc, result.get("score", 0.0))364                ret.append(pair)365        return ret366 367    def max_marginal_relevance_search(368        self,369        query: str,370        k: int = 4,371        fetch_k: int = 20,372        lambda_mult: float = 0.5,373        param: Optional[dict] = None,374        expr: Optional[str] = None,375        **kwargs: Any,376    ) -> List[Document]:377        """Perform a search and return results that are reordered by MMR."""378        embedding = self.embedding_func.embed_query(query)379        return self._max_marginal_relevance_search(380            embedding=embedding,381            k=k,382            fetch_k=fetch_k,383            lambda_mult=lambda_mult,384            param=param,385            expr=expr,386            **kwargs,387        )388 389    def _max_marginal_relevance_search(390        self,391        embedding: list[float],392        k: int = 4,393        fetch_k: int = 20,394        lambda_mult: float = 0.5,395        param: Optional[dict] = None,396        expr: Optional[str] = None,397        **kwargs: Any,398    ) -> List[Document]:399        """Perform a search and return results that are reordered by MMR."""400        ef = 10 if param is None else param.get("ef", 10)401        anns = self.mochowtable.AnnSearch(402            vector_field=self.field_vector,403            vector_floats=[float(num) for num in embedding],404            params=self.mochowtable.HNSWSearchParams(ef=ef, limit=k),405            filter=expr,406        )407        res = self.table.search(anns=anns, retrieve_vector=True)408 409        # Organize results.410        documents: List[Document] = []411        ordered_result_embeddings = []412        rows = [[item] for item in res.rows]413        if rows is None or len(rows) == 0:414            return documents415        for row in rows:416            for result in row:417                row_data = result.get("row", {})418                meta = row_data.get(self.field_metadata)419                if meta is not None:420                    meta = json.loads(meta)421                doc = Document(422                    page_content=row_data.get(self.field_text), metadata=meta423                )424                documents.append(doc)425                ordered_result_embeddings.append(row_data.get(self.field_vector))426        # Get the new order of results.427        new_ordering = maximal_marginal_relevance(428            np.array(embedding), ordered_result_embeddings, k=k, lambda_mult=lambda_mult429        )430        # Reorder the values and return.431        ret = []432        for x in new_ordering:433            # Function can return -1 index434            if x == -1:435                break436            else:437                ret.append(documents[x])438        return ret439 
codekingpro/portable-devtools · Team Ai