Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
lancedb.py696 linesDownload Raw Back to vectorstores
1from __future__ import annotations2 3import base644import os5import uuid6import warnings7from typing import Any, Callable, Dict, Iterable, List, Optional, Type8 9import numpy as np10from langchain_core.documents import Document11from langchain_core.embeddings import Embeddings12from langchain_core.utils import guard_import13from langchain_core.vectorstores import VectorStore14 15from langchain_community.vectorstores.utils import maximal_marginal_relevance16 17DEFAULT_K = 4  # Number of Documents to return.18 19 20def import_lancedb() -> Any:21    """Import lancedb package."""22    return guard_import("lancedb")23 24 25def to_lance_filter(filter: Dict[str, str]) -> str:26    """Converts a dict filter to a LanceDB filter string."""27    return " AND ".join([f"{k} = '{v}'" for k, v in filter.items()])28 29 30class LanceDB(VectorStore):31    """`LanceDB` vector store.32 33    To use, you should have ``lancedb`` python package installed.34    You can install it with ``pip install lancedb``.35 36    Args:37        connection: LanceDB connection to use. If not provided, a new connection38                    will be created.39        embedding: Embedding to use for the vectorstore.40        vector_key: Key to use for the vector in the database. Defaults to ``vector``.41        id_key: Key to use for the id in the database. Defaults to ``id``.42        text_key: Key to use for the text in the database. Defaults to ``text``.43        table_name: Name of the table to use. Defaults to ``vectorstore``.44        api_key: API key to use for LanceDB cloud database.45        region: Region to use for LanceDB cloud database.46        mode: Mode to use for adding data to the table. Valid values are47              ``append`` and ``overwrite``. Defaults to ``overwrite``.48 49 50 51    Example:52        .. code-block:: python53            vectorstore = LanceDB(uri='/lancedb', embedding_function)54            vectorstore.add_texts(['text1', 'text2'])55            result = vectorstore.similarity_search('text1')56    """57 58    def __init__(59        self,60        connection: Optional[Any] = None,61        embedding: Optional[Embeddings] = None,62        uri: Optional[str] = "/tmp/lancedb",63        vector_key: Optional[str] = "vector",64        id_key: Optional[str] = "id",65        text_key: Optional[str] = "text",66        table_name: Optional[str] = "vectorstore",67        api_key: Optional[str] = None,68        region: Optional[str] = None,69        mode: Optional[str] = "overwrite",70        table: Optional[Any] = None,71        distance: Optional[str] = "l2",72        reranker: Optional[Any] = None,73        relevance_score_fn: Optional[Callable[[float], float]] = None,74        limit: int = DEFAULT_K,75    ):76        """Initialize with Lance DB vectorstore"""77        lancedb = guard_import("lancedb")78        lancedb.remote.table = guard_import("lancedb.remote.table")79        self._embedding = embedding80        self._vector_key = vector_key81        self._id_key = id_key82        self._text_key = text_key83        self.api_key = api_key or os.getenv("LANCE_API_KEY") if api_key != "" else None84        self.region = region85        self.mode = mode86        self.distance = distance87        self.override_relevance_score_fn = relevance_score_fn88        self.limit = limit89        self._fts_index = None90 91        if isinstance(reranker, lancedb.rerankers.Reranker):92            self._reranker = reranker93        elif reranker is None:94            self._reranker = None95        else:96            raise ValueError(97                "`reranker` has to be a lancedb.rerankers.Reranker object."98            )99 100        if isinstance(uri, str) and self.api_key is None:101            if uri.startswith("db://"):102                raise ValueError("API key is required for LanceDB cloud.")103 104        if self._embedding is None:105            raise ValueError("embedding object should be provided")106 107        if isinstance(connection, lancedb.db.LanceDBConnection):108            self._connection = connection109        elif isinstance(connection, (str, lancedb.db.LanceTable)):110            raise ValueError(111                "`connection` has to be a lancedb.db.LanceDBConnection object.\112                `lancedb.db.LanceTable` is deprecated."113            )114        else:115            if self.api_key is None:116                self._connection = lancedb.connect(uri)117            else:118                if isinstance(uri, str):119                    if uri.startswith("db://"):120                        self._connection = lancedb.connect(121                            uri, api_key=self.api_key, region=self.region122                        )123                    else:124                        self._connection = lancedb.connect(uri)125                        warnings.warn(126                            "api key provided with local uri.\127                            The data will be stored locally"128                        )129        if table is not None:130            try:131                assert isinstance(132                    table, (lancedb.db.LanceTable, lancedb.remote.table.RemoteTable)133                )134                self._table = table135                self._table_name = (136                    table.name if hasattr(table, "name") else "remote_table"137                )138            except AssertionError:139                raise ValueError(140                    """`table` has to be a lancedb.db.LanceTable or 141                    lancedb.remote.table.RemoteTable object."""142                )143        else:144            self._table = self.get_table(table_name, set_default=True)145 146    def results_to_docs(self, results: Any, score: bool = False) -> Any:147        columns = results.schema.names148 149        if "_distance" in columns:150            score_col = "_distance"151        elif "_relevance_score" in columns:152            score_col = "_relevance_score"153        else:154            score_col = None155        # Check if 'metadata' is in the columns156        has_metadata = "metadata" in columns157 158        if score_col is None or not score:159            return [160                Document(161                    page_content=results[self._text_key][idx].as_py(),162                    metadata=results["metadata"][idx].as_py() if has_metadata else {},163                )164                for idx in range(len(results))165            ]166        elif score_col and score:167            return [168                (169                    Document(170                        page_content=results[self._text_key][idx].as_py(),171                        metadata=results["metadata"][idx].as_py()172                        if has_metadata173                        else {},174                    ),175                    results[score_col][idx].as_py(),176                )177                for idx in range(len(results))178            ]179 180    @property181    def embeddings(self) -> Optional[Embeddings]:182        return self._embedding183 184    def add_texts(185        self,186        texts: Iterable[str],187        metadatas: Optional[List[dict]] = None,188        ids: Optional[List[str]] = None,189        **kwargs: Any,190    ) -> List[str]:191        """Turn texts into embedding and add it to the database192 193        Args:194            texts: Iterable of strings to add to the vectorstore.195            metadatas: Optional list of metadatas associated with the texts.196            ids: Optional list of ids to associate with the texts.197            ids: Optional list of ids to associate with the texts.198 199        Returns:200            List of ids of the added texts.201        """202        docs = []203        ids = ids or [str(uuid.uuid4()) for _ in texts]204        embeddings = self._embedding.embed_documents(list(texts))  # type: ignore[union-attr]205        for idx, text in enumerate(texts):206            embedding = embeddings[idx]207            metadata = metadatas[idx] if metadatas else {"id": ids[idx]}208            docs.append(209                {210                    self._vector_key: embedding,211                    self._id_key: ids[idx],212                    self._text_key: text,213                    "metadata": metadata,214                }215            )216 217        tbl = self.get_table()218 219        if tbl is None:220            tbl = self._connection.create_table(self._table_name, data=docs)221            self._table = tbl222        else:223            if self.api_key is None:224                tbl.add(docs, mode=self.mode)225            else:226                tbl.add(docs)227 228        self._fts_index = None229 230        return ids231 232    def get_table(233        self, name: Optional[str] = None, set_default: Optional[bool] = False234    ) -> Any:235        """236        Fetches a table object from the database.237 238        Args:239            name (str, optional): The name of the table to fetch. Defaults to None240                                    and fetches current table object.241            set_default (bool, optional): Sets fetched table as the default table.242                                        Defaults to False.243 244        Returns:245            Any: The fetched table object.246 247        Raises:248            ValueError: If the specified table is not found in the database.249 250        """251        if name is not None:252            if set_default:253                self._table_name = name254                _name = self._table_name255            else:256                _name = name257        else:258            _name = self._table_name259 260        try:261            return self._connection.open_table(_name)262        except Exception:263            return None264 265    def create_index(266        self,267        col_name: Optional[str] = None,268        vector_col: Optional[str] = None,269        num_partitions: Optional[int] = 256,270        num_sub_vectors: Optional[int] = 96,271        index_cache_size: Optional[int] = None,272        metric: Optional[str] = "L2",273        name: Optional[str] = None,274    ) -> None:275        """276        Create a scalar(for non-vector cols) or a vector index on a table.277        Make sure your vector column has enough data before creating an index on it.278 279        Args:280            vector_col: Provide if you want to create index on a vector column.281            col_name: Provide if you want to create index on a non-vector column.282            metric: Provide the metric to use for vector index. Defaults to 'L2'283                    choice of metrics: 'L2', 'dot', 'cosine'284            num_partitions: Number of partitions to use for the index. Defaults to 256.285            num_sub_vectors: Number of sub-vectors to use for the index. Defaults to 96.286            index_cache_size: Size of the index cache. Defaults to None.287            name: Name of the table to create index on. Defaults to None.288 289        Returns:290            None291        """292        tbl = self.get_table(name)293 294        if vector_col:295            tbl.create_index(296                metric=metric,297                vector_column_name=vector_col,298                num_partitions=num_partitions,299                num_sub_vectors=num_sub_vectors,300                index_cache_size=index_cache_size,301            )302        elif col_name:303            tbl.create_scalar_index(col_name)304        else:305            raise ValueError("Provide either vector_col or col_name")306 307    def encode_image(self, uri: str) -> str:308        """Get base64 string from image URI."""309        with open(uri, "rb") as image_file:310            return base64.b64encode(image_file.read()).decode("utf-8")311 312    def add_images(313        self,314        uris: List[str],315        metadatas: Optional[List[dict]] = None,316        ids: Optional[List[str]] = None,317        **kwargs: Any,318    ) -> List[str]:319        """Run more images through the embeddings and add to the vectorstore.320 321        Args:322            uris List[str]: File path to the image.323            metadatas (Optional[List[dict]], optional): Optional list of metadatas.324            ids (Optional[List[str]], optional): Optional list of IDs.325 326        Returns:327            List[str]: List of IDs of the added images.328        """329        tbl = self.get_table()330 331        # Map from uris to b64 encoded strings332        b64_texts = [self.encode_image(uri=uri) for uri in uris]333        # Populate IDs334        if ids is None:335            ids = [str(uuid.uuid4()) for _ in uris]336        embeddings = None337        # Set embeddings338        if self._embedding is not None and hasattr(self._embedding, "embed_image"):339            embeddings = self._embedding.embed_image(uris=uris)340        else:341            raise ValueError(342                "embedding object should be provided and must have embed_image method."343            )344 345        data = []346        for idx, emb in enumerate(embeddings):347            metadata = metadatas[idx] if metadatas else {"id": ids[idx]}348            data.append(349                {350                    self._vector_key: emb,351                    self._id_key: ids[idx],352                    self._text_key: b64_texts[idx],353                    "metadata": metadata,354                }355            )356        if tbl is None:357            tbl = self._connection.create_table(self._table_name, data=data)358            self._table = tbl359        else:360            tbl.add(data)361 362        return ids363 364    def _query(365        self,366        query: Any,367        k: Optional[int] = None,368        filter: Optional[Any] = None,369        name: Optional[str] = None,370        **kwargs: Any,371    ) -> Any:372        if k is None:373            k = self.limit374        tbl = self.get_table(name)375        if isinstance(filter, dict):376            filter = to_lance_filter(filter)377 378        prefilter = kwargs.get("prefilter", False)379        query_type = kwargs.get("query_type", "vector")380 381        if metrics := kwargs.get("metrics"):382            lance_query = (383                tbl.search(query=query, vector_column_name=self._vector_key)384                .limit(k)385                .metric(metrics)386                .where(filter, prefilter=prefilter)387            )388        else:389            lance_query = (390                tbl.search(query=query, vector_column_name=self._vector_key)391                .limit(k)392                .where(filter, prefilter=prefilter)393            )394        if query_type == "hybrid" and self._reranker is not None:395            lance_query.rerank(reranker=self._reranker)396 397        docs = lance_query.to_arrow()398        if len(docs) == 0:399            warnings.warn("No results found for the query.")400        return docs401 402    def _select_relevance_score_fn(self) -> Callable[[float], float]:403        """404        The 'correct' relevance function405        may differ depending on a few things, including:406        - the distance / similarity metric used by the VectorStore407        - the scale of your embeddings (OpenAI's are unit normed. Many others are not!)408        - embedding dimensionality409        - etc.410        """411        if self.override_relevance_score_fn:412            return self.override_relevance_score_fn413 414        if self.distance == "cosine":415            return self._cosine_relevance_score_fn416        elif self.distance == "l2":417            return self._euclidean_relevance_score_fn418        elif self.distance == "ip":419            return self._max_inner_product_relevance_score_fn420        else:421            raise ValueError(422                "No supported normalization function"423                f" for distance metric of type: {self.distance}."424                "Consider providing relevance_score_fn to Chroma constructor."425            )426 427    def similarity_search_by_vector(428        self,429        embedding: List[float],430        k: Optional[int] = None,431        filter: Optional[Dict[str, str]] = None,432        name: Optional[str] = None,433        **kwargs: Any,434    ) -> Any:435        """436        Return documents most similar to the query vector.437        """438        if k is None:439            k = self.limit440 441        res = self._query(embedding, k, filter=filter, name=name, **kwargs)442        return self.results_to_docs(res, score=kwargs.pop("score", False))443 444    def similarity_search_by_vector_with_relevance_scores(445        self,446        embedding: List[float],447        k: Optional[int] = None,448        filter: Optional[Dict[str, str]] = None,449        name: Optional[str] = None,450        **kwargs: Any,451    ) -> Any:452        """453        Return documents most similar to the query vector with relevance scores.454        """455        if k is None:456            k = self.limit457 458        relevance_score_fn = self._select_relevance_score_fn()459        docs_and_scores = self.similarity_search_by_vector(460            embedding, k, score=True, **kwargs461        )462        return [463            (doc, relevance_score_fn(float(score))) for doc, score in docs_and_scores464        ]465 466    def similarity_search_with_score(467        self,468        query: str,469        k: Optional[int] = None,470        filter: Optional[Dict[str, str]] = None,471        **kwargs: Any,472    ) -> Any:473        """Return documents most similar to the query with relevance scores."""474        if k is None:475            k = self.limit476 477        score = kwargs.get("score", True)478        name = kwargs.get("name", None)479        query_type = kwargs.get("query_type", "vector")480 481        if self._embedding is None:482            raise ValueError("search needs an emmbedding function to be specified.")483 484        if query_type == "fts" or query_type == "hybrid":485            if self.api_key is None and self._fts_index is None:486                tbl = self.get_table(name)487                self._fts_index = tbl.create_fts_index(self._text_key, replace=True)488 489                if query_type == "hybrid":490                    embedding = self._embedding.embed_query(query)491                    _query = (embedding, query)492                else:493                    _query = query  # type: ignore[assignment]494 495                res = self._query(_query, k, filter=filter, name=name, **kwargs)496                return self.results_to_docs(res, score=score)497            else:498                raise NotImplementedError(499                    "Full text/ Hybrid search is not supported in LanceDB Cloud yet."500                )501        else:502            embedding = self._embedding.embed_query(query)503            res = self._query(embedding, k, filter=filter, **kwargs)504            return self.results_to_docs(res, score=score)505 506    def similarity_search(507        self,508        query: str,509        k: Optional[int] = None,510        name: Optional[str] = None,511        filter: Optional[Any] = None,512        fts: Optional[bool] = False,513        **kwargs: Any,514    ) -> List[Document]:515        """Return documents most similar to the query516 517        Args:518            query: String to query the vectorstore with.519            k: Number of documents to return.520            filter (Optional[Dict]): Optional filter arguments521                sql_filter(Optional[string]): SQL filter to apply to the query.522                prefilter(Optional[bool]): Whether to apply the filter prior523                                             to the vector search.524        Raises:525            ValueError: If the specified table is not found in the database.526 527        Returns:528            List of documents most similar to the query.529        """530        res = self.similarity_search_with_score(531            query=query, k=k, name=name, filter=filter, fts=fts, score=False, **kwargs532        )533        return res534 535    def max_marginal_relevance_search(536        self,537        query: str,538        k: Optional[int] = None,539        fetch_k: int = 20,540        lambda_mult: float = 0.5,541        filter: Optional[Dict[str, str]] = None,542        **kwargs: Any,543    ) -> List[Document]:544        """Return docs selected using the maximal marginal relevance.545        Maximal marginal relevance optimizes for similarity to query AND diversity546        among selected documents.547 548        Args:549            query: Text to look up documents similar to.550            k: Number of Documents to return. Defaults to 4.551            fetch_k: Number of Documents to fetch to pass to MMR algorithm.552            lambda_mult: Number between 0 and 1 that determines the degree553                        of diversity among the results with 0 corresponding554                        to maximum diversity and 1 to minimum diversity.555                        Defaults to 0.5.556            filter (Optional[Dict[str, str]]): Filter by metadata. Defaults to None.557 558        Returns:559            List of Documents selected by maximal marginal relevance.560        """561        if k is None:562            k = self.limit563 564        if self._embedding is None:565            raise ValueError(566                "For MMR search, you must specify an embedding function oncreation."567            )568 569        embedding = self._embedding.embed_query(query)570        docs = self.max_marginal_relevance_search_by_vector(571            embedding,572            k,573            fetch_k,574            lambda_mult=lambda_mult,575            filter=filter,576        )577        return docs578 579    def max_marginal_relevance_search_by_vector(580        self,581        embedding: List[float],582        k: Optional[int] = None,583        fetch_k: int = 20,584        lambda_mult: float = 0.5,585        filter: Optional[Dict[str, str]] = None,586        **kwargs: Any,587    ) -> List[Document]:588        """Return docs selected using the maximal marginal relevance.589        Maximal marginal relevance optimizes for similarity to query AND diversity590        among selected documents.591 592        Args:593            embedding: Embedding to look up documents similar to.594            k: Number of Documents to return. Defaults to 4.595            fetch_k: Number of Documents to fetch to pass to MMR algorithm.596            lambda_mult: Number between 0 and 1 that determines the degree597                        of diversity among the results with 0 corresponding598                        to maximum diversity and 1 to minimum diversity.599                        Defaults to 0.5.600            filter (Optional[Dict[str, str]]): Filter by metadata. Defaults to None.601 602        Returns:603            List of Documents selected by maximal marginal relevance.604        """605 606        results = self._query(607            query=embedding,608            k=fetch_k,609            filter=filter,610            **kwargs,611        )612        mmr_selected = maximal_marginal_relevance(613            np.array(embedding, dtype=np.float32),614            results["vector"].to_pylist(),615            k=k or self.limit,616            lambda_mult=lambda_mult,617        )618 619        candidates = self.results_to_docs(results)620 621        selected_results = [r for i, r in enumerate(candidates) if i in mmr_selected]622        return selected_results623 624    @classmethod625    def from_texts(626        cls: Type[LanceDB],627        texts: List[str],628        embedding: Embeddings,629        metadatas: Optional[List[dict]] = None,630        connection: Optional[Any] = None,631        vector_key: Optional[str] = "vector",632        id_key: Optional[str] = "id",633        text_key: Optional[str] = "text",634        table_name: Optional[str] = "vectorstore",635        api_key: Optional[str] = None,636        region: Optional[str] = None,637        mode: Optional[str] = "overwrite",638        distance: Optional[str] = "l2",639        reranker: Optional[Any] = None,640        relevance_score_fn: Optional[Callable[[float], float]] = None,641        **kwargs: Any,642    ) -> LanceDB:643        instance = LanceDB(644            connection=connection,645            embedding=embedding,646            vector_key=vector_key,647            id_key=id_key,648            text_key=text_key,649            table_name=table_name,650            api_key=api_key,651            region=region,652            mode=mode,653            distance=distance,654            reranker=reranker,655            relevance_score_fn=relevance_score_fn,656            **kwargs,657        )658        instance.add_texts(texts, metadatas=metadatas)659 660        return instance661 662    def delete(663        self,664        ids: Optional[List[str]] = None,665        delete_all: Optional[bool] = None,666        filter: Optional[str] = None,667        drop_columns: Optional[List[str]] = None,668        name: Optional[str] = None,669        **kwargs: Any,670    ) -> None:671        """672        Allows deleting rows by filtering, by ids or drop columns from the table.673 674        Args:675            filter: Provide a string SQL expression -  "{col} {operation} {value}".676            ids: Provide list of ids to delete from the table.677            drop_columns: Provide list of columns to drop from the table.678            delete_all: If True, delete all rows from the table.679        """680        tbl = self.get_table(name)681        if filter:682            tbl.delete(filter)683        elif ids:684            tbl.delete(f"{self._id_key} in ('{{}}')".format(",".join(ids)))685        elif drop_columns:686            if self.api_key is not None:687                raise NotImplementedError(688                    "Column operations currently not supported in LanceDB Cloud."689                )690            else:691                tbl.drop_columns(drop_columns)692        elif delete_all:693            tbl.delete("true")694        else:695            raise ValueError("Provide either filter, ids, drop_columns or delete_all")696 
codekingpro/portable-devtools · Team Ai