codekingpro/portable-devtools
114k
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 