Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
rocksetdb.py422 linesDownload Raw Back to vectorstores
1from __future__ import annotations2 3import logging4from copy import deepcopy5from enum import Enum6from typing import Any, Iterable, List, Optional, Tuple7 8import numpy as np9from langchain_core.documents import Document10from langchain_core.embeddings import Embeddings11from langchain_core.runnables import run_in_executor12from langchain_core.vectorstores import VectorStore13 14from langchain_community.vectorstores.utils import maximal_marginal_relevance15 16logger = logging.getLogger(__name__)17 18 19class Rockset(VectorStore):20    """`Rockset` vector store.21 22    To use, you should have the `rockset` python package installed. Note that to use23    this, the collection being used must already exist in your Rockset instance.24    You must also ensure you use a Rockset ingest transformation to apply25    `VECTOR_ENFORCE` on the column being used to store `embedding_key` in the26    collection.27    See: https://rockset.com/blog/introducing-vector-search-on-rockset/ for more details28 29    Everything below assumes `commons` Rockset workspace.30 31    Example:32        .. code-block:: python33 34            from langchain_community.vectorstores import Rockset35            from langchain_community.embeddings.openai import OpenAIEmbeddings36            import rockset37 38            # Make sure you use the right host (region) for your Rockset instance39            # and APIKEY has both read-write access to your collection.40 41            rs = rockset.RocksetClient(host=rockset.Regions.use1a1, api_key="***")42            collection_name = "langchain_demo"43            embeddings = OpenAIEmbeddings()44            vectorstore = Rockset(rs, collection_name, embeddings,45                "description", "description_embedding")46 47    """48 49    def __init__(50        self,51        client: Any,52        embeddings: Embeddings,53        collection_name: str,54        text_key: str,55        embedding_key: str,56        workspace: str = "commons",57    ):58        """Initialize with Rockset client.59        Args:60            client: Rockset client object61            collection: Rockset collection to insert docs / query62            embeddings: Langchain Embeddings object to use to generate63                        embedding for given text.64            text_key: column in Rockset collection to use to store the text65            embedding_key: column in Rockset collection to use to store the embedding.66                           Note: We must apply `VECTOR_ENFORCE()` on this column via67                           Rockset ingest transformation.68 69        """70        try:71            from rockset import RocksetClient72        except ImportError:73            raise ImportError(74                "Could not import rockset client python package. "75                "Please install it with `pip install rockset`."76            )77 78        if not isinstance(client, RocksetClient):79            raise ValueError(80                f"client should be an instance of rockset.RocksetClient, "81                f"got {type(client)}"82            )83        # TODO: check that `collection_name` exists in rockset. Create if not.84        self._client = client85        self._collection_name = collection_name86        self._embeddings = embeddings87        self._text_key = text_key88        self._embedding_key = embedding_key89        self._workspace = workspace90 91        try:92            self._client.set_application("langchain")93        except AttributeError:94            # ignore95            pass96 97    @property98    def embeddings(self) -> Embeddings:99        return self._embeddings100 101    def add_texts(102        self,103        texts: Iterable[str],104        metadatas: Optional[List[dict]] = None,105        ids: Optional[List[str]] = None,106        batch_size: int = 32,107        **kwargs: Any,108    ) -> List[str]:109        """Run more texts through the embeddings and add to the vectorstore110 111                Args:112            texts: Iterable of strings to add to the vectorstore.113            metadatas: Optional list of metadatas associated with the texts.114            ids: Optional list of ids to associate with the texts.115            batch_size: Send documents in batches to rockset.116 117        Returns:118            List of ids from adding the texts into the vectorstore.119 120        """121        batch: list[dict] = []122        stored_ids = []123 124        for i, text in enumerate(texts):125            if len(batch) == batch_size:126                stored_ids += self._write_documents_to_rockset(batch)127                batch = []128            doc = {}129            if metadatas and len(metadatas) > i:130                doc = deepcopy(metadatas[i])131            if ids and len(ids) > i:132                doc["_id"] = ids[i]133            doc[self._text_key] = text134            doc[self._embedding_key] = self._embeddings.embed_query(text)135            batch.append(doc)136        if len(batch) > 0:137            stored_ids += self._write_documents_to_rockset(batch)138            batch = []139        return stored_ids140 141    @classmethod142    def from_texts(143        cls,144        texts: List[str],145        embedding: Embeddings,146        metadatas: Optional[List[dict]] = None,147        client: Any = None,148        collection_name: str = "",149        text_key: str = "",150        embedding_key: str = "",151        ids: Optional[List[str]] = None,152        batch_size: int = 32,153        **kwargs: Any,154    ) -> Rockset:155        """Create Rockset wrapper with existing texts.156        This is intended as a quicker way to get started.157        """158 159        # Sanitize inputs160        assert client is not None, "Rockset Client cannot be None"161        assert collection_name, "Collection name cannot be empty"162        assert text_key, "Text key name cannot be empty"163        assert embedding_key, "Embedding key cannot be empty"164 165        rockset = cls(client, embedding, collection_name, text_key, embedding_key)166        rockset.add_texts(texts, metadatas, ids, batch_size)167        return rockset168 169    # Rockset supports these vector distance functions.170    class DistanceFunction(Enum):171        COSINE_SIM = "COSINE_SIM"172        EUCLIDEAN_DIST = "EUCLIDEAN_DIST"173        DOT_PRODUCT = "DOT_PRODUCT"174 175        # how to sort results for "similarity"176        def order_by(self) -> str:177            if self.value == "EUCLIDEAN_DIST":178                return "ASC"179            return "DESC"180 181    def similarity_search_with_relevance_scores(182        self,183        query: str,184        k: int = 4,185        distance_func: DistanceFunction = DistanceFunction.COSINE_SIM,186        where_str: Optional[str] = None,187        **kwargs: Any,188    ) -> List[Tuple[Document, float]]:189        """Perform a similarity search with Rockset190 191        Args:192            query (str): Text to look up documents similar to.193            distance_func (DistanceFunction): how to compute distance between two194                vectors in Rockset.195            k (int, optional): Top K neighbors to retrieve. Defaults to 4.196            where_str (Optional[str], optional): Metadata filters supplied as a197                SQL `where` condition string. Defaults to None.198                eg. "price<=70.0 AND brand='Nintendo'"199 200            NOTE: Please do not let end-user to fill this and always be aware201                  of SQL injection.202 203        Returns:204            List[Tuple[Document, float]]: List of documents with their relevance score205        """206        return self.similarity_search_by_vector_with_relevance_scores(207            self._embeddings.embed_query(query),208            k,209            distance_func,210            where_str,211            **kwargs,212        )213 214    def similarity_search(215        self,216        query: str,217        k: int = 4,218        distance_func: DistanceFunction = DistanceFunction.COSINE_SIM,219        where_str: Optional[str] = None,220        **kwargs: Any,221    ) -> List[Document]:222        """Same as `similarity_search_with_relevance_scores` but223        doesn't return the scores.224        """225        return self.similarity_search_by_vector(226            self._embeddings.embed_query(query),227            k,228            distance_func,229            where_str,230            **kwargs,231        )232 233    def similarity_search_by_vector(234        self,235        embedding: List[float],236        k: int = 4,237        distance_func: DistanceFunction = DistanceFunction.COSINE_SIM,238        where_str: Optional[str] = None,239        **kwargs: Any,240    ) -> List[Document]:241        """Accepts a query_embedding (vector), and returns documents with242        similar embeddings."""243 244        docs_and_scores = self.similarity_search_by_vector_with_relevance_scores(245            embedding, k, distance_func, where_str, **kwargs246        )247        return [doc for doc, _ in docs_and_scores]248 249    def similarity_search_by_vector_with_relevance_scores(250        self,251        embedding: List[float],252        k: int = 4,253        distance_func: DistanceFunction = DistanceFunction.COSINE_SIM,254        where_str: Optional[str] = None,255        **kwargs: Any,256    ) -> List[Tuple[Document, float]]:257        """Accepts a query_embedding (vector), and returns documents with258        similar embeddings along with their relevance scores."""259 260        exclude_embeddings = True261        if "exclude_embeddings" in kwargs:262            exclude_embeddings = kwargs["exclude_embeddings"]263        q_str = self._build_query_sql(264            embedding, distance_func, k, where_str, exclude_embeddings265        )266        try:267            query_response = self._client.Queries.query(sql={"query": q_str})268        except Exception as e:269            logger.error("Exception when querying Rockset: %s\n", e)270            return []271        finalResult: list[Tuple[Document, float]] = []272        for document in query_response.results:273            metadata = {}274            assert isinstance(document, dict), (275                "document should be of type `dict[str,Any]`. But found: `{}`".format(276                    type(document)277                )278            )279            for k, v in document.items():280                if k == self._text_key:281                    assert isinstance(v, str), (282                        "page content stored in column `{}` must be of type `str`. "283                        "But found: `{}`"284                    ).format(self._text_key, type(v))285                    page_content = v286                elif k == "dist":287                    assert isinstance(v, float), (288                        "Computed distance between vectors must of type `float`. "289                        "But found {}"290                    ).format(type(v))291                    score = v292                elif k not in ["_id", "_event_time", "_meta"]:293                    # These columns are populated by Rockset when documents are294                    # inserted. No need to return them in metadata dict.295                    metadata[k] = v296            finalResult.append(297                (298                    Document(page_content=page_content, metadata=metadata),299                    score,300                )301            )302        return finalResult303 304    def max_marginal_relevance_search(305        self,306        query: str,307        k: int = 4,308        fetch_k: int = 20,309        lambda_mult: float = 0.5,310        *,311        where_str: Optional[str] = None,312        **kwargs: Any,313    ) -> List[Document]:314        """Return docs selected using the maximal marginal relevance.315 316        Maximal marginal relevance optimizes for similarity to query AND diversity317        among selected documents.318 319        Args:320            query: Text to look up documents similar to.321            k: Number of Documents to return. Defaults to 4.322            fetch_k: Number of Documents to fetch to pass to MMR algorithm.323            distance_func (DistanceFunction): how to compute distance between two324                vectors in Rockset.325            lambda_mult: Number between 0 and 1 that determines the degree326                        of diversity among the results with 0 corresponding327                        to maximum diversity and 1 to minimum diversity.328                        Defaults to 0.5.329            where_str: where clause for the sql query330        Returns:331            List of Documents selected by maximal marginal relevance.332        """333        query_embedding = self._embeddings.embed_query(query)334        initial_docs = self.similarity_search_by_vector(335            query_embedding,336            k=fetch_k,337            where_str=where_str,338            exclude_embeddings=False,339            **kwargs,340        )341 342        embeddings = [doc.metadata[self._embedding_key] for doc in initial_docs]343 344        selected_indices = maximal_marginal_relevance(345            np.array(query_embedding),346            embeddings,347            lambda_mult=lambda_mult,348            k=k,349        )350 351        # remove embeddings key before returning for cleanup to be consistent with352        #   other search functions353        for i in selected_indices:354            del initial_docs[i].metadata[self._embedding_key]355 356        return [initial_docs[i] for i in selected_indices]357 358    # Helper functions359 360    def _build_query_sql(361        self,362        query_embedding: List[float],363        distance_func: DistanceFunction,364        k: int = 4,365        where_str: Optional[str] = None,366        exclude_embeddings: bool = True,367    ) -> str:368        """Builds Rockset SQL query to query similar vectors to query_vector"""369 370        q_embedding_str = ",".join(map(str, query_embedding))371        distance_str = f"""{distance_func.value}({self._embedding_key}, \372[{q_embedding_str}]) as dist"""373        where_str = f"WHERE {where_str}\n" if where_str else ""374        select_embedding = (375            f" EXCEPT({self._embedding_key})," if exclude_embeddings else ","376        )377        return f"""\378SELECT *{select_embedding} {distance_str}379FROM {self._workspace}.{self._collection_name}380{where_str}\381ORDER BY dist {distance_func.order_by()}382LIMIT {str(k)}383"""384 385    def _write_documents_to_rockset(self, batch: List[dict]) -> List[str]:386        add_doc_res = self._client.Documents.add_documents(387            collection=self._collection_name, data=batch, workspace=self._workspace388        )389        return [doc_status._id for doc_status in add_doc_res.data]390 391    def delete_texts(self, ids: List[str]) -> None:392        """Delete a list of docs from the Rockset collection"""393        try:394            from rockset.models import DeleteDocumentsRequestData395        except ImportError:396            raise ImportError(397                "Could not import rockset client python package. "398                "Please install it with `pip install rockset`."399            )400 401        self._client.Documents.delete_documents(402            collection=self._collection_name,403            data=[DeleteDocumentsRequestData(id=i) for i in ids],404            workspace=self._workspace,405        )406 407    def delete(self, ids: Optional[List[str]] = None, **kwargs: Any) -> Optional[bool]:408        try:409            if ids is None:410                ids = []411            self.delete_texts(ids)412        except Exception as e:413            logger.error("Exception when deleting docs from Rockset: %s\n", e)414            return False415 416        return True417 418    async def adelete(419        self, ids: Optional[List[str]] = None, **kwargs: Any420    ) -> Optional[bool]:421        return await run_in_executor(None, self.delete, ids, **kwargs)422 
codekingpro/portable-devtools · Team Ai