Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
milvus.py1094 linesDownload Raw Back to vectorstores
1from __future__ import annotations2 3import logging4from typing import TYPE_CHECKING, Any, Iterable, List, Optional, Tuple, Union5from uuid import uuid46 7import numpy as np8from langchain_core._api.deprecation import deprecated9from langchain_core.documents import Document10from langchain_core.embeddings import Embeddings11from langchain_core.vectorstores import VectorStore12 13from langchain_community.vectorstores.utils import maximal_marginal_relevance14 15if TYPE_CHECKING:16    from pymilvus.orm.mutation import MutationResult17 18logger = logging.getLogger(__name__)19 20DEFAULT_MILVUS_CONNECTION = {21    "host": "localhost",22    "port": "19530",23    "user": "",24    "password": "",25    "secure": False,26}27 28 29@deprecated(30    since="0.2.0",31    removal="1.0",32    alternative_import="langchain_milvus.MilvusVectorStore",33)34class Milvus(VectorStore):35    """`Milvus` vector store.36 37    You need to install `pymilvus` and run Milvus.38 39    See the following documentation for how to run a Milvus instance:40    https://milvus.io/docs/install_standalone-docker.md41 42    If looking for a hosted Milvus, take a look at this documentation:43    https://zilliz.com/cloud and make use of the Zilliz vectorstore found in44    this project.45 46    IF USING L2/IP metric, IT IS HIGHLY SUGGESTED TO NORMALIZE YOUR DATA.47 48    Args:49        embedding_function (Embeddings): Function used to embed the text.50        collection_name (str): Which Milvus collection to use. Defaults to51            "LangChainCollection".52        collection_description (str): The description of the collection. Defaults to53            "".54        collection_properties (Optional[dict[str, any]]): The collection properties.55            Defaults to None.56            If set, will override collection existing properties.57            For example: {"collection.ttl.seconds": 60}.58        connection_args (Optional[dict[str, any]]): The connection args used for59            this class comes in the form of a dict.60        consistency_level (str): The consistency level to use for a collection.61            Defaults to "Session".62        index_params (Optional[dict]): Which index params to use. Defaults to63            HNSW/AUTOINDEX depending on service.64        search_params (Optional[dict]): Which search params to use. Defaults to65            default of index.66        drop_old (Optional[bool]): Whether to drop the current collection. Defaults67            to False.68        auto_id (bool): Whether to enable auto id for primary key. Defaults to False.69            If False, you needs to provide text ids (string less than 65535 bytes).70            If True, Milvus will generate unique integers as primary keys.71        primary_field (str): Name of the primary key field. Defaults to "pk".72        text_field (str): Name of the text field. Defaults to "text".73        vector_field (str): Name of the vector field. Defaults to "vector".74        metadata_field (str): Name of the metadata field. Defaults to None.75            When metadata_field is specified,76            the document's metadata will store as json.77 78    The connection args used for this class comes in the form of a dict,79    here are a few of the options:80        address (str): The actual address of Milvus81            instance. Example address: "localhost:19530"82        uri (str): The uri of Milvus instance. Example uri:83            "http://randomwebsite:19530",84            "tcp:foobarsite:19530",85            "https://ok.s3.south.com:19530".86        host (str): The host of Milvus instance. Default at "localhost",87            PyMilvus will fill in the default host if only port is provided.88        port (str/int): The port of Milvus instance. Default at 19530, PyMilvus89            will fill in the default port if only host is provided.90        user (str): Use which user to connect to Milvus instance. If user and91            password are provided, we will add related header in every RPC call.92        password (str): Required when user is provided. The password93            corresponding to the user.94        secure (bool): Default is false. If set to true, tls will be enabled.95        client_key_path (str): If use tls two-way authentication, need to96            write the client.key path.97        client_pem_path (str): If use tls two-way authentication, need to98            write the client.pem path.99        ca_pem_path (str): If use tls two-way authentication, need to write100            the ca.pem path.101        server_pem_path (str): If use tls one-way authentication, need to102            write the server.pem path.103        server_name (str): If use tls, need to write the common name.104 105    Example:106        .. code-block:: python107 108        from langchain_community.vectorstores import Milvus109        from langchain_community.embeddings import OpenAIEmbeddings110 111        embedding = OpenAIEmbeddings()112        # Connect to a milvus instance on localhost113        milvus_store = Milvus(114            embedding_function = Embeddings,115            collection_name = "LangChainCollection",116            drop_old = True,117            auto_id = True118        )119 120    Raises:121        ValueError: If the pymilvus python package is not installed.122    """123 124    def __init__(125        self,126        embedding_function: Embeddings,127        collection_name: str = "LangChainCollection",128        collection_description: str = "",129        collection_properties: Optional[dict[str, Any]] = None,130        connection_args: Optional[dict[str, Any]] = None,131        consistency_level: str = "Session",132        index_params: Optional[dict] = None,133        search_params: Optional[dict] = None,134        drop_old: Optional[bool] = False,135        auto_id: bool = False,136        *,137        primary_field: str = "pk",138        text_field: str = "text",139        vector_field: str = "vector",140        metadata_field: Optional[str] = None,141        partition_key_field: Optional[str] = None,142        partition_names: Optional[list] = None,143        replica_number: int = 1,144        timeout: Optional[float] = None,145        num_shards: Optional[int] = None,146    ):147        """Initialize the Milvus vector store."""148        try:149            from pymilvus import Collection, utility150        except ImportError:151            raise ImportError(152                "Could not import pymilvus python package. "153                "Please install it with `pip install pymilvus`."154            )155 156        # Default search params when one is not provided.157        self.default_search_params = {158            "IVF_FLAT": {"metric_type": "L2", "params": {"nprobe": 10}},159            "IVF_SQ8": {"metric_type": "L2", "params": {"nprobe": 10}},160            "IVF_PQ": {"metric_type": "L2", "params": {"nprobe": 10}},161            "HNSW": {"metric_type": "L2", "params": {"ef": 10}},162            "RHNSW_FLAT": {"metric_type": "L2", "params": {"ef": 10}},163            "RHNSW_SQ": {"metric_type": "L2", "params": {"ef": 10}},164            "RHNSW_PQ": {"metric_type": "L2", "params": {"ef": 10}},165            "IVF_HNSW": {"metric_type": "L2", "params": {"nprobe": 10, "ef": 10}},166            "ANNOY": {"metric_type": "L2", "params": {"search_k": 10}},167            "SCANN": {"metric_type": "L2", "params": {"search_k": 10}},168            "AUTOINDEX": {"metric_type": "L2", "params": {}},169            "GPU_CAGRA": {170                "metric_type": "L2",171                "params": {172                    "itopk_size": 128,173                    "search_width": 4,174                    "min_iterations": 0,175                    "max_iterations": 0,176                    "team_size": 0,177                },178            },179            "GPU_IVF_FLAT": {"metric_type": "L2", "params": {"nprobe": 10}},180            "GPU_IVF_PQ": {"metric_type": "L2", "params": {"nprobe": 10}},181        }182 183        self.embedding_func = embedding_function184        self.collection_name = collection_name185        self.collection_description = collection_description186        self.collection_properties = collection_properties187        self.index_params = index_params188        self.search_params = search_params189        self.consistency_level = consistency_level190        self.auto_id = auto_id191 192        # In order for a collection to be compatible, pk needs to be varchar193        self._primary_field = primary_field194        # In order for compatibility, the text field will need to be called "text"195        self._text_field = text_field196        # In order for compatibility, the vector field needs to be called "vector"197        self._vector_field = vector_field198        self._metadata_field = metadata_field199        self._partition_key_field = partition_key_field200        self.fields: list[str] = []201        self.partition_names = partition_names202        self.replica_number = replica_number203        self.timeout = timeout204        self.num_shards = num_shards205 206        # Create the connection to the server207        if connection_args is None:208            connection_args = DEFAULT_MILVUS_CONNECTION209        self.alias = self._create_connection_alias(connection_args)210        self.col: Optional[Collection] = None211 212        # Grab the existing collection if it exists213        if utility.has_collection(self.collection_name, using=self.alias):214            self.col = Collection(215                self.collection_name,216                using=self.alias,217            )218            if self.collection_properties is not None:219                self.col.set_properties(self.collection_properties)220        # If need to drop old, drop it221        if drop_old and isinstance(self.col, Collection):222            self.col.drop()223            self.col = None224 225        # Initialize the vector store226        self._init(227            partition_names=partition_names,228            replica_number=replica_number,229            timeout=timeout,230        )231 232    @property233    def embeddings(self) -> Embeddings:234        return self.embedding_func235 236    def _create_connection_alias(self, connection_args: dict) -> str:237        """Create the connection to the Milvus server."""238        from pymilvus import MilvusException, connections239 240        # Grab the connection arguments that are used for checking existing connection241        host: Optional[str] = connection_args.get("host", None)242        port: Optional[Union[str, int]] = connection_args.get("port", None)243        address: Optional[str] = connection_args.get("address", None)244        uri: Optional[str] = connection_args.get("uri", None)245        user = connection_args.get("user", None)246 247        # Order of use is host/port, uri, address248        if host is not None and port is not None:249            given_address = str(host) + ":" + str(port)250        elif uri is not None:251            if uri.startswith("https://"):252                given_address = uri.split("https://")[1]253            elif uri.startswith("http://"):254                given_address = uri.split("http://")[1]255            else:256                logger.error("Invalid Milvus URI: %s", uri)257                raise ValueError("Invalid Milvus URI: %s", uri)258        elif address is not None:259            given_address = address260        else:261            given_address = None262            logger.debug("Missing standard address type for reuse attempt")263 264        # User defaults to empty string when getting connection info265        if user is not None:266            tmp_user = user267        else:268            tmp_user = ""269 270        # If a valid address was given, then check if a connection exists271        if given_address is not None:272            for con in connections.list_connections():273                addr = connections.get_connection_addr(con[0])274                if (275                    con[1]276                    and ("address" in addr)277                    and (addr["address"] == given_address)278                    and ("user" in addr)279                    and (addr["user"] == tmp_user)280                ):281                    logger.debug("Using previous connection: %s", con[0])282                    return con[0]283 284        # Generate a new connection if one doesn't exist285        alias = uuid4().hex286        try:287            connections.connect(alias=alias, **connection_args)288            logger.debug("Created new connection using: %s", alias)289            return alias290        except MilvusException as e:291            logger.error("Failed to create new connection using: %s", alias)292            raise e293 294    def _init(295        self,296        embeddings: Optional[list] = None,297        metadatas: Optional[list[dict]] = None,298        partition_names: Optional[list] = None,299        replica_number: int = 1,300        timeout: Optional[float] = None,301    ) -> None:302        if embeddings is not None:303            self._create_collection(embeddings, metadatas)304        self._extract_fields()305        self._create_index()306        self._create_search_params()307        self._load(308            partition_names=partition_names,309            replica_number=replica_number,310            timeout=timeout,311        )312 313    def _create_collection(314        self, embeddings: list, metadatas: Optional[list[dict]] = None315    ) -> None:316        from pymilvus import (317            Collection,318            CollectionSchema,319            DataType,320            FieldSchema,321            MilvusException,322        )323        from pymilvus.orm.types import infer_dtype_bydata324 325        # Determine embedding dim326        dim = len(embeddings[0])327        fields = []328        if self._metadata_field is not None:329            fields.append(FieldSchema(self._metadata_field, DataType.JSON))330        else:331            # Determine metadata schema332            if metadatas:333                # Create FieldSchema for each entry in metadata.334                for key, value in metadatas[0].items():335                    # Infer the corresponding datatype of the metadata336                    dtype = infer_dtype_bydata(value)337                    # Datatype isn't compatible338                    if dtype == DataType.UNKNOWN or dtype == DataType.NONE:339                        logger.error(340                            (341                                "Failure to create collection, "342                                "unrecognized dtype for key: %s"343                            ),344                            key,345                        )346                        raise ValueError(f"Unrecognized datatype for {key}.")347                    # Dataype is a string/varchar equivalent348                    elif dtype == DataType.VARCHAR:349                        fields.append(350                            FieldSchema(key, DataType.VARCHAR, max_length=65_535)351                        )352                    else:353                        fields.append(FieldSchema(key, dtype))354 355        # Create the text field356        fields.append(357            FieldSchema(self._text_field, DataType.VARCHAR, max_length=65_535)358        )359        # Create the primary key field360        if self.auto_id:361            fields.append(362                FieldSchema(363                    self._primary_field, DataType.INT64, is_primary=True, auto_id=True364                )365            )366        else:367            fields.append(368                FieldSchema(369                    self._primary_field,370                    DataType.VARCHAR,371                    is_primary=True,372                    auto_id=False,373                    max_length=65_535,374                )375            )376        # Create the vector field, supports binary or float vectors377        fields.append(378            FieldSchema(self._vector_field, infer_dtype_bydata(embeddings[0]), dim=dim)379        )380 381        # Create the schema for the collection382        schema = CollectionSchema(383            fields,384            description=self.collection_description,385            partition_key_field=self._partition_key_field,386        )387 388        # Create the collection389        try:390            if self.num_shards is not None:391                # Issue with defaults:392                # https://github.com/milvus-io/pymilvus/blob/59bf5e811ad56e20946559317fed855330758d9c/pymilvus/client/prepare.py#L82-L85393                self.col = Collection(394                    name=self.collection_name,395                    schema=schema,396                    consistency_level=self.consistency_level,397                    using=self.alias,398                    num_shards=self.num_shards,399                )400            else:401                self.col = Collection(402                    name=self.collection_name,403                    schema=schema,404                    consistency_level=self.consistency_level,405                    using=self.alias,406                )407            # Set the collection properties if they exist408            if self.collection_properties is not None:409                self.col.set_properties(self.collection_properties)410        except MilvusException as e:411            logger.error(412                "Failed to create collection: %s error: %s", self.collection_name, e413            )414            raise e415 416    def _extract_fields(self) -> None:417        """Grab the existing fields from the Collection"""418        from pymilvus import Collection419 420        if isinstance(self.col, Collection):421            schema = self.col.schema422            for x in schema.fields:423                self.fields.append(x.name)424 425    def _get_index(self) -> Optional[dict[str, Any]]:426        """Return the vector index information if it exists"""427        from pymilvus import Collection428 429        if isinstance(self.col, Collection):430            for x in self.col.indexes:431                if x.field_name == self._vector_field:432                    return x.to_dict()433        return None434 435    def _create_index(self) -> None:436        """Create a index on the collection"""437        from pymilvus import Collection, MilvusException438 439        if isinstance(self.col, Collection) and self._get_index() is None:440            try:441                # If no index params, use a default HNSW based one442                if self.index_params is None:443                    self.index_params = {444                        "metric_type": "L2",445                        "index_type": "HNSW",446                        "params": {"M": 8, "efConstruction": 64},447                    }448 449                try:450                    self.col.create_index(451                        self._vector_field,452                        index_params=self.index_params,453                        using=self.alias,454                    )455 456                # If default did not work, most likely on Zilliz Cloud457                except MilvusException:458                    # Use AUTOINDEX based index459                    self.index_params = {460                        "metric_type": "L2",461                        "index_type": "AUTOINDEX",462                        "params": {},463                    }464                    self.col.create_index(465                        self._vector_field,466                        index_params=self.index_params,467                        using=self.alias,468                    )469                logger.debug(470                    "Successfully created an index on collection: %s",471                    self.collection_name,472                )473 474            except MilvusException as e:475                logger.error(476                    "Failed to create an index on collection: %s", self.collection_name477                )478                raise e479 480    def _create_search_params(self) -> None:481        """Generate search params based on the current index type"""482        from pymilvus import Collection483 484        if isinstance(self.col, Collection) and self.search_params is None:485            index = self._get_index()486            if index is not None:487                index_type: str = index["index_param"]["index_type"]488                metric_type: str = index["index_param"]["metric_type"]489                self.search_params = self.default_search_params[index_type]490                self.search_params["metric_type"] = metric_type491 492    def _load(493        self,494        partition_names: Optional[list] = None,495        replica_number: int = 1,496        timeout: Optional[float] = None,497    ) -> None:498        """Load the collection if available."""499        from pymilvus import Collection, utility500        from pymilvus.client.types import LoadState501 502        timeout = self.timeout or timeout503        if (504            isinstance(self.col, Collection)505            and self._get_index() is not None506            and utility.load_state(self.collection_name, using=self.alias)507            == LoadState.NotLoad508        ):509            self.col.load(510                partition_names=partition_names,511                replica_number=replica_number,512                timeout=timeout,513            )514 515    def add_texts(516        self,517        texts: Iterable[str],518        metadatas: Optional[List[dict]] = None,519        timeout: Optional[float] = None,520        batch_size: int = 1000,521        *,522        ids: Optional[List[str]] = None,523        **kwargs: Any,524    ) -> List[str]:525        """Insert text data into Milvus.526 527        Inserting data when the collection has not be made yet will result528        in creating a new Collection. The data of the first entity decides529        the schema of the new collection, the dim is extracted from the first530        embedding and the columns are decided by the first metadata dict.531        Metadata keys will need to be present for all inserted values. At532        the moment there is no None equivalent in Milvus.533 534        Args:535            texts (Iterable[str]): The texts to embed, it is assumed536                that they all fit in memory.537            metadatas (Optional[List[dict]]): Metadata dicts attached to each of538                the texts. Defaults to None.539            should be less than 65535 bytes. Required and work when auto_id is False.540            timeout (Optional[float]): Timeout for each batch insert. Defaults541                to None.542            batch_size (int, optional): Batch size to use for insertion.543                Defaults to 1000.544            ids (Optional[List[str]]): List of text ids. The length of each item545 546        Raises:547            MilvusException: Failure to add texts548 549        Returns:550            List[str]: The resulting keys for each inserted element.551        """552        from pymilvus import Collection, MilvusException553 554        texts = list(texts)555        if not self.auto_id:556            assert isinstance(ids, list), (557                "A list of valid ids are required when auto_id is False."558            )559            assert len(set(ids)) == len(texts), (560                "Different lengths of texts and unique ids are provided."561            )562            assert all(len(x.encode()) <= 65_535 for x in ids), (563                "Each id should be a string less than 65535 bytes."564            )565 566        try:567            embeddings = self.embedding_func.embed_documents(texts)568        except NotImplementedError:569            embeddings = [self.embedding_func.embed_query(x) for x in texts]570 571        if len(embeddings) == 0:572            logger.debug("Nothing to insert, skipping.")573            return []574 575        # If the collection hasn't been initialized yet, perform all steps to do so576        if not isinstance(self.col, Collection):577            kwargs = {"embeddings": embeddings, "metadatas": metadatas}578            if self.partition_names:579                kwargs["partition_names"] = self.partition_names580            if self.replica_number:581                kwargs["replica_number"] = self.replica_number582            if self.timeout:583                kwargs["timeout"] = self.timeout584            self._init(**kwargs)585 586        # Dict to hold all insert columns587        insert_dict: dict[str, list] = {588            self._text_field: texts,589            self._vector_field: embeddings,590        }591 592        if not self.auto_id:593            insert_dict[self._primary_field] = ids  # type: ignore[assignment]594 595        if self._metadata_field is not None:596            for d in metadatas:  # type: ignore[union-attr]597                insert_dict.setdefault(self._metadata_field, []).append(d)598        else:599            # Collect the metadata into the insert dict.600            if metadatas is not None:601                for d in metadatas:602                    for key, value in d.items():603                        keys = (604                            [x for x in self.fields if x != self._primary_field]605                            if self.auto_id606                            else [x for x in self.fields]607                        )608                        if key in keys:609                            insert_dict.setdefault(key, []).append(value)610 611        # Total insert count612        vectors: list = insert_dict[self._vector_field]613        total_count = len(vectors)614 615        pks: list[str] = []616 617        assert isinstance(self.col, Collection)618        for i in range(0, total_count, batch_size):619            # Grab end index620            end = min(i + batch_size, total_count)621            # Convert dict to list of lists batch for insertion622            insert_list = [623                insert_dict[x][i:end] for x in self.fields if x in insert_dict624            ]625            # Insert into the collection.626            try:627                res: Collection628                timeout = self.timeout or timeout629                res = self.col.insert(insert_list, timeout=timeout, **kwargs)630                pks.extend(res.primary_keys)631            except MilvusException as e:632                logger.error(633                    "Failed to insert batch starting at entity: %s/%s", i, total_count634                )635                raise e636        return pks637 638    def similarity_search(639        self,640        query: str,641        k: int = 4,642        param: Optional[dict] = None,643        expr: Optional[str] = None,644        timeout: Optional[float] = None,645        **kwargs: Any,646    ) -> List[Document]:647        """Perform a similarity search against the query string.648 649        Args:650            query (str): The text to search.651            k (int, optional): How many results to return. Defaults to 4.652            param (dict, optional): The search params for the index type.653                Defaults to None.654            expr (str, optional): Filtering expression. Defaults to None.655            timeout (int, optional): How long to wait before timeout error.656                Defaults to None.657            kwargs: Collection.search() keyword arguments.658 659        Returns:660            List[Document]: Document results for search.661        """662        if self.col is None:663            logger.debug("No existing collection to search.")664            return []665        timeout = self.timeout or timeout666        res = self.similarity_search_with_score(667            query=query, k=k, param=param, expr=expr, timeout=timeout, **kwargs668        )669        return [doc for doc, _ in res]670 671    def similarity_search_by_vector(672        self,673        embedding: List[float],674        k: int = 4,675        param: Optional[dict] = None,676        expr: Optional[str] = None,677        timeout: Optional[float] = None,678        **kwargs: Any,679    ) -> List[Document]:680        """Perform a similarity search against the query string.681 682        Args:683            embedding (List[float]): The embedding vector to search.684            k (int, optional): How many results to return. Defaults to 4.685            param (dict, optional): The search params for the index type.686                Defaults to None.687            expr (str, optional): Filtering expression. Defaults to None.688            timeout (int, optional): How long to wait before timeout error.689                Defaults to None.690            kwargs: Collection.search() keyword arguments.691 692        Returns:693            List[Document]: Document results for search.694        """695        if self.col is None:696            logger.debug("No existing collection to search.")697            return []698        timeout = self.timeout or timeout699        res = self.similarity_search_with_score_by_vector(700            embedding=embedding, k=k, param=param, expr=expr, timeout=timeout, **kwargs701        )702        return [doc for doc, _ in res]703 704    def similarity_search_with_score(705        self,706        query: str,707        k: int = 4,708        param: Optional[dict] = None,709        expr: Optional[str] = None,710        timeout: Optional[float] = None,711        **kwargs: Any,712    ) -> List[Tuple[Document, float]]:713        """Perform a search on a query string and return results with score.714 715        For more information about the search parameters, take a look at the pymilvus716        documentation found here:717        https://milvus.io/api-reference/pymilvus/v2.2.6/Collection/search().md718 719        Args:720            query (str): The text being searched.721            k (int, optional): The amount of results to return. Defaults to 4.722            param (dict): The search params for the specified index.723                Defaults to None.724            expr (str, optional): Filtering expression. Defaults to None.725            timeout (float, optional): How long to wait before timeout error.726                Defaults to None.727            kwargs: Collection.search() keyword arguments.728 729        Returns:730            List[float], List[Tuple[Document, any, any]]:731        """732        if self.col is None:733            logger.debug("No existing collection to search.")734            return []735 736        # Embed the query text.737        embedding = self.embedding_func.embed_query(query)738        timeout = self.timeout or timeout739        res = self.similarity_search_with_score_by_vector(740            embedding=embedding, k=k, param=param, expr=expr, timeout=timeout, **kwargs741        )742        return res743 744    def similarity_search_with_score_by_vector(745        self,746        embedding: List[float],747        k: int = 4,748        param: Optional[dict] = None,749        expr: Optional[str] = None,750        timeout: Optional[float] = None,751        **kwargs: Any,752    ) -> List[Tuple[Document, float]]:753        """Perform a search on a query string and return results with score.754 755        For more information about the search parameters, take a look at the pymilvus756        documentation found here:757        https://milvus.io/api-reference/pymilvus/v2.2.6/Collection/search().md758 759        Args:760            embedding (List[float]): The embedding vector being searched.761            k (int, optional): The amount of results to return. Defaults to 4.762            param (dict): The search params for the specified index.763                Defaults to None.764            expr (str, optional): Filtering expression. Defaults to None.765            timeout (float, optional): How long to wait before timeout error.766                Defaults to None.767            kwargs: Collection.search() keyword arguments.768 769        Returns:770            List[Tuple[Document, float]]: Result doc and score.771        """772        if self.col is None:773            logger.debug("No existing collection to search.")774            return []775 776        if param is None:777            param = self.search_params778 779        # Determine result metadata fields with PK.780        output_fields = self.fields[:]781        output_fields.remove(self._vector_field)782        timeout = self.timeout or timeout783        # Perform the search.784        res = self.col.search(785            data=[embedding],786            anns_field=self._vector_field,787            param=param,788            limit=k,789            expr=expr,790            output_fields=output_fields,791            timeout=timeout,792            **kwargs,793        )794        # Organize results.795        ret = []796        for result in res[0]:797            data = {x: result.entity.get(x) for x in output_fields}798            doc = self._parse_document(data)799            pair = (doc, result.score)800            ret.append(pair)801 802        return ret803 804    def max_marginal_relevance_search(805        self,806        query: str,807        k: int = 4,808        fetch_k: int = 20,809        lambda_mult: float = 0.5,810        param: Optional[dict] = None,811        expr: Optional[str] = None,812        timeout: Optional[float] = None,813        **kwargs: Any,814    ) -> List[Document]:815        """Perform a search and return results that are reordered by MMR.816 817        Args:818            query (str): The text being searched.819            k (int, optional): How many results to give. Defaults to 4.820            fetch_k (int, optional): Total results to select k from.821                Defaults to 20.822            lambda_mult: Number between 0 and 1 that determines the degree823                        of diversity among the results with 0 corresponding824                        to maximum diversity and 1 to minimum diversity.825                        Defaults to 0.5826            param (dict, optional): The search params for the specified index.827                Defaults to None.828            expr (str, optional): Filtering expression. Defaults to None.829            timeout (float, optional): How long to wait before timeout error.830                Defaults to None.831            kwargs: Collection.search() keyword arguments.832 833 834        Returns:835            List[Document]: Document results for search.836        """837        if self.col is None:838            logger.debug("No existing collection to search.")839            return []840 841        embedding = self.embedding_func.embed_query(query)842        timeout = self.timeout or timeout843        return self.max_marginal_relevance_search_by_vector(844            embedding=embedding,845            k=k,846            fetch_k=fetch_k,847            lambda_mult=lambda_mult,848            param=param,849            expr=expr,850            timeout=timeout,851            **kwargs,852        )853 854    def max_marginal_relevance_search_by_vector(855        self,856        embedding: list[float],857        k: int = 4,858        fetch_k: int = 20,859        lambda_mult: float = 0.5,860        param: Optional[dict] = None,861        expr: Optional[str] = None,862        timeout: Optional[float] = None,863        **kwargs: Any,864    ) -> List[Document]:865        """Perform a search and return results that are reordered by MMR.866 867        Args:868            embedding (str): The embedding vector being searched.869            k (int, optional): How many results to give. Defaults to 4.870            fetch_k (int, optional): Total results to select k from.871                Defaults to 20.872            lambda_mult: Number between 0 and 1 that determines the degree873                        of diversity among the results with 0 corresponding874                        to maximum diversity and 1 to minimum diversity.875                        Defaults to 0.5876            param (dict, optional): The search params for the specified index.877                Defaults to None.878            expr (str, optional): Filtering expression. Defaults to None.879            timeout (float, optional): How long to wait before timeout error.880                Defaults to None.881            kwargs: Collection.search() keyword arguments.882 883        Returns:884            List[Document]: Document results for search.885        """886        if self.col is None:887            logger.debug("No existing collection to search.")888            return []889 890        if param is None:891            param = self.search_params892 893        # Determine result metadata fields.894        output_fields = self.fields[:]895        output_fields.remove(self._vector_field)896        timeout = self.timeout or timeout897        # Perform the search.898        res = self.col.search(899            data=[embedding],900            anns_field=self._vector_field,901            param=param,902            limit=fetch_k,903            expr=expr,904            output_fields=output_fields,905            timeout=timeout,906            **kwargs,907        )908        # Organize results.909        ids = []910        documents = []911        scores = []912        for result in res[0]:913            data = {x: result.entity.get(x) for x in output_fields}914            doc = self._parse_document(data)915            documents.append(doc)916            scores.append(result.score)917            ids.append(result.id)918 919        vectors = self.col.query(920            expr=f"{self._primary_field} in {ids}",921            output_fields=[self._primary_field, self._vector_field],922            timeout=timeout,923        )924        # Reorganize the results from query to match search order.925        vectors = {x[self._primary_field]: x[self._vector_field] for x in vectors}926 927        ordered_result_embeddings = [vectors[x] for x in ids]928 929        # Get the new order of results.930        new_ordering = maximal_marginal_relevance(931            np.array(embedding), ordered_result_embeddings, k=k, lambda_mult=lambda_mult932        )933 934        # Reorder the values and return.935        ret = []936        for x in new_ordering:937            # Function can return -1 index938            if x == -1:939                break940            else:941                ret.append(documents[x])942        return ret943 944    def delete(945        self, ids: Optional[List[str]] = None, expr: Optional[str] = None, **kwargs: Any946    ) -> MutationResult:947        """Delete by vector ID or boolean expression.948        Refer to [Milvus documentation](https://milvus.io/docs/delete_data.md)949        for notes and examples of expressions.950 951        Args:952            ids: List of ids to delete.953            expr: Boolean expression that specifies the entities to delete.954            kwargs: Other parameters in Milvus delete api.955        """956        if isinstance(ids, list) and len(ids) > 0:957            if expr is not None:958                logger.warning(959                    "Both ids and expr are provided. Ignore expr and delete by ids."960                )961            expr = f"{self._primary_field} in {ids}"962        else:963            assert isinstance(expr, str), (964                "Either ids list or expr string must be provided."965            )966        return self.col.delete(expr=expr, **kwargs)  # type: ignore[union-attr]967 968    @classmethod969    def from_texts(970        cls,971        texts: List[str],972        embedding: Embeddings,973        metadatas: Optional[List[dict]] = None,974        collection_name: str = "LangChainCollection",975        connection_args: dict[str, Any] = DEFAULT_MILVUS_CONNECTION,976        consistency_level: str = "Session",977        index_params: Optional[dict] = None,978        search_params: Optional[dict] = None,979        drop_old: bool = False,980        *,981        ids: Optional[List[str]] = None,982        **kwargs: Any,983    ) -> Milvus:984        """Create a Milvus collection, indexes it with HNSW, and insert data.985 986        Args:987            texts (List[str]): Text data.988            embedding (Embeddings): Embedding function.989            metadatas (Optional[List[dict]]): Metadata for each text if it exists.990                Defaults to None.991            collection_name (str, optional): Collection name to use. Defaults to992                "LangChainCollection".993            connection_args (dict[str, Any], optional): Connection args to use. Defaults994                to DEFAULT_MILVUS_CONNECTION.995            consistency_level (str, optional): Which consistency level to use. Defaults996                to "Session".997            index_params (Optional[dict], optional): Which index_params to use. Defaults998                to None.999            search_params (Optional[dict], optional): Which search params to use.1000                Defaults to None.1001            drop_old (Optional[bool], optional): Whether to drop the collection with1002                that name if it exists. Defaults to False.1003            ids (Optional[List[str]]): List of text ids. Defaults to None.1004 1005        Returns:1006            Milvus: Milvus Vector Store1007        """1008        if isinstance(ids, list) and len(ids) > 0:1009            auto_id = False1010        else:1011            auto_id = True1012 1013        vector_db = cls(1014            embedding_function=embedding,1015            collection_name=collection_name,1016            connection_args=connection_args,1017            consistency_level=consistency_level,1018            index_params=index_params,1019            search_params=search_params,1020            drop_old=drop_old,1021            auto_id=auto_id,1022            **kwargs,1023        )1024        vector_db.add_texts(texts=texts, metadatas=metadatas, ids=ids)1025        return vector_db1026 1027    def _parse_document(self, data: dict) -> Document:1028        return Document(1029            page_content=data.pop(self._text_field),1030            metadata=data.pop(self._metadata_field) if self._metadata_field else data,1031        )1032 1033    def get_pks(self, expr: str, **kwargs: Any) -> List[int] | None:1034        """Get primary keys with expression1035 1036        Args:1037            expr: Expression - E.g: "id in [1, 2]", or "title LIKE 'Abc%'"1038 1039        Returns:1040            List[int]: List of IDs (Primary Keys)1041        """1042 1043        from pymilvus import MilvusException1044 1045        if self.col is None:1046            logger.debug("No existing collection to get pk.")1047            return None1048 1049        try:1050            query_result = self.col.query(1051                expr=expr, output_fields=[self._primary_field]1052            )1053        except MilvusException as exc:1054            logger.error("Failed to get ids: %s error: %s", self.collection_name, exc)1055            raise exc1056        pks = [item.get(self._primary_field) for item in query_result]1057        return pks1058 1059    def upsert(1060        self,1061        ids: Optional[List[str]] = None,1062        documents: List[Document] | None = None,1063        **kwargs: Any,1064    ) -> List[str] | None:1065        """Update/Insert documents to the vectorstore.1066 1067        Args:1068            ids: IDs to update - Let's call get_pks to get ids with expression \n1069            documents (List[Document]): Documents to add to the vectorstore.1070 1071        Returns:1072            List[str]: IDs of the added texts.1073        """1074 1075        from pymilvus import MilvusException1076 1077        if documents is None or len(documents) == 0:1078            logger.debug("No documents to upsert.")1079            return None1080 1081        if ids is not None and len(ids):1082            kwargs["ids"] = ids1083            try:1084                self.delete(ids=ids)1085            except MilvusException:1086                pass1087        try:1088            return self.add_documents(documents=documents, **kwargs)1089        except MilvusException as exc:1090            logger.error(1091                "Failed to upsert entities: %s error: %s", self.collection_name, exc1092            )1093            raise exc1094 
codekingpro/portable-devtools · Team Ai