Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
sklearn.py355 linesDownload Raw Back to vectorstores
1"""Wrapper around scikit-learn NearestNeighbors implementation.2 3The vector store can be persisted in json, bson or parquet format.4"""5 6import json7import math8import os9from abc import ABC, abstractmethod10from typing import Any, Dict, Iterable, List, Literal, Optional, Tuple, Type11from uuid import uuid412 13from langchain_core.documents import Document14from langchain_core.embeddings import Embeddings15from langchain_core.utils import guard_import16from langchain_core.vectorstores import VectorStore17 18from langchain_community.vectorstores.utils import maximal_marginal_relevance19 20DEFAULT_K = 4  # Number of Documents to return.21DEFAULT_FETCH_K = 20  # Number of Documents to initially fetch during MMR search.22 23 24class BaseSerializer(ABC):25    """Base class for serializing data."""26 27    def __init__(self, persist_path: str) -> None:28        self.persist_path = persist_path29 30    @classmethod31    @abstractmethod32    def extension(cls) -> str:33        """The file extension suggested by this serializer (without dot)."""34 35    @abstractmethod36    def save(self, data: Any) -> None:37        """Saves the data to the persist_path"""38 39    @abstractmethod40    def load(self) -> Any:41        """Loads the data from the persist_path"""42 43 44class JsonSerializer(BaseSerializer):45    """Serialize data in JSON using the json package from python standard library."""46 47    @classmethod48    def extension(cls) -> str:49        return "json"50 51    def save(self, data: Any) -> None:52        with open(self.persist_path, "w") as fp:53            json.dump(data, fp)54 55    def load(self) -> Any:56        with open(self.persist_path, "r") as fp:57            return json.load(fp)58 59 60class BsonSerializer(BaseSerializer):61    """Serialize data in Binary JSON using the `bson` python package."""62 63    def __init__(self, persist_path: str) -> None:64        super().__init__(persist_path)65        self.bson = guard_import("bson")66 67    @classmethod68    def extension(cls) -> str:69        return "bson"70 71    def save(self, data: Any) -> None:72        with open(self.persist_path, "wb") as fp:73            fp.write(self.bson.dumps(data))74 75    def load(self) -> Any:76        with open(self.persist_path, "rb") as fp:77            return self.bson.loads(fp.read())78 79 80class ParquetSerializer(BaseSerializer):81    """Serialize data in `Apache Parquet` format using the `pyarrow` package."""82 83    def __init__(self, persist_path: str) -> None:84        super().__init__(persist_path)85        self.pd = guard_import("pandas")86        self.pa = guard_import("pyarrow")87        self.pq = guard_import("pyarrow.parquet")88 89    @classmethod90    def extension(cls) -> str:91        return "parquet"92 93    def save(self, data: Any) -> None:94        df = self.pd.DataFrame(data)95        table = self.pa.Table.from_pandas(df)96        if os.path.exists(self.persist_path):97            backup_path = str(self.persist_path) + "-backup"98            os.rename(self.persist_path, backup_path)99            try:100                self.pq.write_table(table, self.persist_path)101            except Exception as exc:102                os.rename(backup_path, self.persist_path)103                raise exc104            else:105                os.remove(backup_path)106        else:107            self.pq.write_table(table, self.persist_path)108 109    def load(self) -> Any:110        table = self.pq.read_table(self.persist_path)111        df = table.to_pandas()112        return {col: series.tolist() for col, series in df.items()}113 114 115SERIALIZER_MAP: Dict[str, Type[BaseSerializer]] = {116    "json": JsonSerializer,117    "bson": BsonSerializer,118    "parquet": ParquetSerializer,119}120 121 122class SKLearnVectorStoreException(RuntimeError):123    """Exception raised by SKLearnVectorStore."""124 125    pass126 127 128class SKLearnVectorStore(VectorStore):129    """Simple in-memory vector store based on the `scikit-learn` library130    `NearestNeighbors`."""131 132    def __init__(133        self,134        embedding: Embeddings,135        *,136        persist_path: Optional[str] = None,137        serializer: Literal["json", "bson", "parquet"] = "json",138        metric: str = "cosine",139        **kwargs: Any,140    ) -> None:141        np = guard_import("numpy")142        sklearn_neighbors = guard_import("sklearn.neighbors", pip_name="scikit-learn")143 144        # non-persistent properties145        self._np = np146        self._neighbors = sklearn_neighbors.NearestNeighbors(metric=metric, **kwargs)147        self._neighbors_fitted = False148        self._embedding_function = embedding149        self._persist_path = persist_path150        self._serializer: Optional[BaseSerializer] = None151        if self._persist_path is not None:152            serializer_cls = SERIALIZER_MAP[serializer]153            self._serializer = serializer_cls(persist_path=self._persist_path)154 155        # data properties156        self._embeddings: List[List[float]] = []157        self._texts: List[str] = []158        self._metadatas: List[dict] = []159        self._ids: List[str] = []160 161        # cache properties162        self._embeddings_np: Any = np.asarray([])163 164        if self._persist_path is not None and os.path.isfile(self._persist_path):165            self._load()166 167    @property168    def embeddings(self) -> Embeddings:169        return self._embedding_function170 171    def persist(self) -> None:172        if self._serializer is None:173            raise SKLearnVectorStoreException(174                "You must specify a persist_path on creation to persist the collection."175            )176        data = {177            "ids": self._ids,178            "texts": self._texts,179            "metadatas": self._metadatas,180            "embeddings": self._embeddings,181        }182        self._serializer.save(data)183 184    def _load(self) -> None:185        if self._serializer is None:186            raise SKLearnVectorStoreException(187                "You must specify a persist_path on creation to load the collection."188            )189        data = self._serializer.load()190        self._embeddings = data["embeddings"]191        self._texts = data["texts"]192        self._metadatas = data["metadatas"]193        self._ids = data["ids"]194        self._update_neighbors()195 196    def add_texts(197        self,198        texts: Iterable[str],199        metadatas: Optional[List[dict]] = None,200        ids: Optional[List[str]] = None,201        **kwargs: Any,202    ) -> List[str]:203        _texts = list(texts)204        _ids = ids or [str(uuid4()) for _ in _texts]205        self._texts.extend(_texts)206        self._embeddings.extend(self._embedding_function.embed_documents(_texts))207        self._metadatas.extend(metadatas or ([{}] * len(_texts)))208        self._ids.extend(_ids)209        self._update_neighbors()210        return _ids211 212    def _update_neighbors(self) -> None:213        if len(self._embeddings) == 0:214            raise SKLearnVectorStoreException(215                "No data was added to SKLearnVectorStore."216            )217        self._embeddings_np = self._np.asarray(self._embeddings)218        self._neighbors.fit(self._embeddings_np)219        self._neighbors_fitted = True220 221    def _similarity_index_search_with_score(222        self, query_embedding: List[float], *, k: int = DEFAULT_K, **kwargs: Any223    ) -> List[Tuple[int, float]]:224        """Search k embeddings similar to the query embedding. Returns a list of225        (index, distance) tuples."""226        if not self._neighbors_fitted:227            raise SKLearnVectorStoreException(228                "No data was added to SKLearnVectorStore."229            )230        neigh_dists, neigh_idxs = self._neighbors.kneighbors(231            [query_embedding], n_neighbors=k232        )233        return list(zip(neigh_idxs[0], neigh_dists[0]))234 235    def similarity_search_with_score(236        self, query: str, *, k: int = DEFAULT_K, **kwargs: Any237    ) -> List[Tuple[Document, float]]:238        query_embedding = self._embedding_function.embed_query(query)239        indices_dists = self._similarity_index_search_with_score(240            query_embedding, k=k, **kwargs241        )242        return [243            (244                Document(245                    page_content=self._texts[idx],246                    metadata={"id": self._ids[idx], **self._metadatas[idx]},247                ),248                dist,249            )250            for idx, dist in indices_dists251        ]252 253    def similarity_search(254        self, query: str, k: int = DEFAULT_K, **kwargs: Any255    ) -> List[Document]:256        docs_scores = self.similarity_search_with_score(query, k=k, **kwargs)257        return [doc for doc, _ in docs_scores]258 259    def _similarity_search_with_relevance_scores(260        self, query: str, k: int = DEFAULT_K, **kwargs: Any261    ) -> List[Tuple[Document, float]]:262        docs_dists = self.similarity_search_with_score(query, k=k, **kwargs)263        docs, dists = zip(*docs_dists)264        scores = [1 / math.exp(dist) for dist in dists]265        return list(zip(list(docs), scores))266 267    def max_marginal_relevance_search_by_vector(268        self,269        embedding: List[float],270        k: int = DEFAULT_K,271        fetch_k: int = DEFAULT_FETCH_K,272        lambda_mult: float = 0.5,273        **kwargs: Any,274    ) -> List[Document]:275        """Return docs selected using the maximal marginal relevance.276        Maximal marginal relevance optimizes for similarity to query AND diversity277        among selected documents.278        Args:279            embedding: Embedding to look up documents similar to.280            k: Number of Documents to return. Defaults to 4.281            fetch_k: Number of Documents to fetch to pass to MMR algorithm.282            lambda_mult: Number between 0 and 1 that determines the degree283                        of diversity among the results with 0 corresponding284                        to maximum diversity and 1 to minimum diversity.285                        Defaults to 0.5.286        Returns:287            List of Documents selected by maximal marginal relevance.288        """289        indices_dists = self._similarity_index_search_with_score(290            embedding, k=fetch_k, **kwargs291        )292        indices, _ = zip(*indices_dists)293        result_embeddings = self._embeddings_np[indices,]294        mmr_selected = maximal_marginal_relevance(295            self._np.array(embedding, dtype=self._np.float32),296            result_embeddings,297            k=k,298            lambda_mult=lambda_mult,299        )300        mmr_indices = [indices[i] for i in mmr_selected]301        return [302            Document(303                page_content=self._texts[idx],304                metadata={"id": self._ids[idx], **self._metadatas[idx]},305            )306            for idx in mmr_indices307        ]308 309    def max_marginal_relevance_search(310        self,311        query: str,312        k: int = DEFAULT_K,313        fetch_k: int = DEFAULT_FETCH_K,314        lambda_mult: float = 0.5,315        **kwargs: Any,316    ) -> List[Document]:317        """Return docs selected using the maximal marginal relevance.318        Maximal marginal relevance optimizes for similarity to query AND diversity319        among selected documents.320        Args:321            query: Text to look up documents similar to.322            k: Number of Documents to return. Defaults to 4.323            fetch_k: Number of Documents to fetch to pass to MMR algorithm.324            lambda_mult: Number between 0 and 1 that determines the degree325                        of diversity among the results with 0 corresponding326                        to maximum diversity and 1 to minimum diversity.327                        Defaults to 0.5.328        Returns:329            List of Documents selected by maximal marginal relevance.330        """331        if self._embedding_function is None:332            raise ValueError(333                "For MMR search, you must specify an embedding function on creation."334            )335 336        embedding = self._embedding_function.embed_query(query)337        docs = self.max_marginal_relevance_search_by_vector(338            embedding, k, fetch_k, lambda_mul=lambda_mult339        )340        return docs341 342    @classmethod343    def from_texts(344        cls,345        texts: List[str],346        embedding: Embeddings,347        metadatas: Optional[List[dict]] = None,348        ids: Optional[List[str]] = None,349        persist_path: Optional[str] = None,350        **kwargs: Any,351    ) -> "SKLearnVectorStore":352        vs = SKLearnVectorStore(embedding, persist_path=persist_path, **kwargs)353        vs.add_texts(texts, metadatas=metadatas, ids=ids)354        return vs355 
codekingpro/portable-devtools · Team Ai