codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import operator4import pickle5import uuid6from pathlib import Path7from typing import Any, Callable, Dict, Iterable, List, Optional, Tuple8 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.docstore.base import AddableMixin, Docstore16from langchain_community.docstore.in_memory import InMemoryDocstore17from langchain_community.vectorstores.utils import DistanceStrategy18 19 20def normalize(x: np.ndarray) -> np.ndarray:21 """Normalize vectors to unit length."""22 x /= np.clip(np.linalg.norm(x, axis=-1, keepdims=True), 1e-12, None)23 return x24 25 26def dependable_scann_import() -> Any:27 """28 Import `scann` if available, otherwise raise error.29 """30 return guard_import("scann")31 32 33class ScaNN(VectorStore):34 """`ScaNN` vector store.35 36 To use, you should have the ``scann`` python package installed.37 38 Example:39 .. code-block:: python40 41 from langchain_community.embeddings import HuggingFaceEmbeddings42 from langchain_community.vectorstores import ScaNN43 44 model_name = "sentence-transformers/all-mpnet-base-v2"45 db = ScaNN.from_texts(46 ['foo', 'bar', 'barz', 'qux'],47 HuggingFaceEmbeddings(model_name=model_name))48 db.similarity_search('foo?', k=1)49 """50 51 def __init__(52 self,53 embedding: Embeddings,54 index: Any,55 docstore: Docstore,56 index_to_docstore_id: Dict[int, str],57 relevance_score_fn: Optional[Callable[[float], float]] = None,58 normalize_L2: bool = False,59 distance_strategy: DistanceStrategy = DistanceStrategy.EUCLIDEAN_DISTANCE,60 scann_config: Optional[str] = None,61 ):62 """Initialize with necessary components."""63 self.embedding = embedding64 self.index = index65 self.docstore = docstore66 self.index_to_docstore_id = index_to_docstore_id67 self.distance_strategy = distance_strategy68 self.override_relevance_score_fn = relevance_score_fn69 self._normalize_L2 = normalize_L270 self._scann_config = scann_config71 72 def __add(73 self,74 texts: Iterable[str],75 embeddings: Iterable[List[float]],76 metadatas: Optional[List[dict]] = None,77 ids: Optional[List[str]] = None,78 **kwargs: Any,79 ) -> List[str]:80 if not isinstance(self.docstore, AddableMixin):81 raise ValueError(82 "If trying to add texts, the underlying docstore should support "83 f"adding items, which {self.docstore} does not"84 )85 raise NotImplementedError("Updates are not available in ScaNN, yet.")86 87 def add_texts(88 self,89 texts: Iterable[str],90 metadatas: Optional[List[dict]] = None,91 ids: Optional[List[str]] = None,92 **kwargs: Any,93 ) -> List[str]:94 """Run more texts through the embeddings and add to the vectorstore.95 96 Args:97 texts: Iterable of strings to add to the vectorstore.98 metadatas: Optional list of metadatas associated with the texts.99 ids: Optional list of unique IDs.100 101 Returns:102 List of ids from adding the texts into the vectorstore.103 """104 # Embed and create the documents.105 embeddings = self.embedding.embed_documents(list(texts))106 return self.__add(texts, embeddings, metadatas=metadatas, ids=ids, **kwargs)107 108 def add_embeddings(109 self,110 text_embeddings: Iterable[Tuple[str, List[float]]],111 metadatas: Optional[List[dict]] = None,112 ids: Optional[List[str]] = None,113 **kwargs: Any,114 ) -> List[str]:115 """Run more texts through the embeddings and add to the vectorstore.116 117 Args:118 text_embeddings: Iterable pairs of string and embedding to119 add to the vectorstore.120 metadatas: Optional list of metadatas associated with the texts.121 ids: Optional list of unique IDs.122 123 Returns:124 List of ids from adding the texts into the vectorstore.125 """126 if not isinstance(self.docstore, AddableMixin):127 raise ValueError(128 "If trying to add texts, the underlying docstore should support "129 f"adding items, which {self.docstore} does not"130 )131 # Embed and create the documents.132 texts, embeddings = zip(*text_embeddings)133 134 return self.__add(texts, embeddings, metadatas=metadatas, ids=ids, **kwargs)135 136 def delete(self, ids: Optional[List[str]] = None, **kwargs: Any) -> Optional[bool]:137 """Delete by vector ID or other criteria.138 139 Args:140 ids: List of ids to delete.141 **kwargs: Other keyword arguments that subclasses might use.142 143 Returns:144 Optional[bool]: True if deletion is successful,145 False otherwise, None if not implemented.146 """147 148 raise NotImplementedError("Deletions are not available in ScaNN, yet.")149 150 def similarity_search_with_score_by_vector(151 self,152 embedding: List[float],153 k: int = 4,154 filter: Optional[Dict[str, Any]] = None,155 fetch_k: int = 20,156 **kwargs: Any,157 ) -> List[Tuple[Document, float]]:158 """Return docs most similar to query.159 160 Args:161 embedding: Embedding vector to look up documents similar to.162 k: Number of Documents to return. Defaults to 4.163 filter (Optional[Dict[str, Any]]): Filter by metadata. Defaults to None.164 fetch_k: (Optional[int]) Number of Documents to fetch before filtering.165 Defaults to 20.166 **kwargs: kwargs to be passed to similarity search. Can include:167 score_threshold: Optional, a floating point value between 0 to 1 to168 filter the resulting set of retrieved docs169 170 Returns:171 List of documents most similar to the query text and L2 distance172 in float for each. Lower score represents more similarity.173 """174 vector = np.array([embedding], dtype=np.float32)175 if self._normalize_L2:176 vector = normalize(vector)177 indices, scores = self.index.search_batched(178 vector, k if filter is None else fetch_k179 )180 docs = []181 for j, i in enumerate(indices[0]):182 if i == -1:183 # This happens when not enough docs are returned.184 continue185 _id = self.index_to_docstore_id[i]186 doc = self.docstore.search(_id)187 if not isinstance(doc, Document):188 raise ValueError(f"Could not find document for id {_id}, got {doc}")189 if filter is not None:190 filter = {191 key: [value] if not isinstance(value, list) else value192 for key, value in filter.items()193 }194 if all(doc.metadata.get(key) in value for key, value in filter.items()):195 docs.append((doc, scores[0][j]))196 else:197 docs.append((doc, scores[0][j]))198 199 score_threshold = kwargs.get("score_threshold")200 if score_threshold is not None:201 cmp = (202 operator.ge203 if self.distance_strategy204 in (DistanceStrategy.MAX_INNER_PRODUCT, DistanceStrategy.JACCARD)205 else operator.le206 )207 docs = [208 (doc, similarity)209 for doc, similarity in docs210 if cmp(similarity, score_threshold)211 ]212 return docs[:k]213 214 def similarity_search_with_score(215 self,216 query: str,217 k: int = 4,218 filter: Optional[Dict[str, Any]] = None,219 fetch_k: int = 20,220 **kwargs: Any,221 ) -> List[Tuple[Document, float]]:222 """Return docs most similar to query.223 224 Args:225 query: Text to look up documents similar to.226 k: Number of Documents to return. Defaults to 4.227 filter (Optional[Dict[str, str]]): Filter by metadata. Defaults to None.228 fetch_k: (Optional[int]) Number of Documents to fetch before filtering.229 Defaults to 20.230 231 Returns:232 List of documents most similar to the query text with233 L2 distance in float. Lower score represents more similarity.234 """235 embedding = self.embedding.embed_query(query)236 docs = self.similarity_search_with_score_by_vector(237 embedding,238 k,239 filter=filter,240 fetch_k=fetch_k,241 **kwargs,242 )243 return docs244 245 def similarity_search_by_vector(246 self,247 embedding: List[float],248 k: int = 4,249 filter: Optional[Dict[str, Any]] = None,250 fetch_k: int = 20,251 **kwargs: Any,252 ) -> List[Document]:253 """Return docs most similar to embedding vector.254 255 Args:256 embedding: Embedding to look up documents similar to.257 k: Number of Documents to return. Defaults to 4.258 filter (Optional[Dict[str, str]]): Filter by metadata. Defaults to None.259 fetch_k: (Optional[int]) Number of Documents to fetch before filtering.260 Defaults to 20.261 262 Returns:263 List of Documents most similar to the embedding.264 """265 docs_and_scores = self.similarity_search_with_score_by_vector(266 embedding,267 k,268 filter=filter,269 fetch_k=fetch_k,270 **kwargs,271 )272 return [doc for doc, _ in docs_and_scores]273 274 def similarity_search(275 self,276 query: str,277 k: int = 4,278 filter: Optional[Dict[str, Any]] = None,279 fetch_k: int = 20,280 **kwargs: Any,281 ) -> List[Document]:282 """Return docs most similar to query.283 284 Args:285 query: Text to look up documents similar to.286 k: Number of Documents to return. Defaults to 4.287 filter: (Optional[Dict[str, str]]): Filter by metadata. Defaults to None.288 fetch_k: (Optional[int]) Number of Documents to fetch before filtering.289 Defaults to 20.290 291 Returns:292 List of Documents most similar to the query.293 """294 docs_and_scores = self.similarity_search_with_score(295 query, k, filter=filter, fetch_k=fetch_k, **kwargs296 )297 return [doc for doc, _ in docs_and_scores]298 299 @classmethod300 def __from(301 cls,302 texts: List[str],303 embeddings: List[List[float]],304 embedding: Embeddings,305 metadatas: Optional[List[dict]] = None,306 ids: Optional[List[str]] = None,307 normalize_L2: bool = False,308 **kwargs: Any,309 ) -> ScaNN:310 scann = guard_import("scann")311 distance_strategy = kwargs.get(312 "distance_strategy", DistanceStrategy.EUCLIDEAN_DISTANCE313 )314 scann_config = kwargs.get("scann_config", None)315 316 vector = np.array(embeddings, dtype=np.float32)317 if normalize_L2:318 vector = normalize(vector)319 if scann_config is not None:320 index = scann.scann_ops_pybind.create_searcher(vector, scann_config)321 else:322 if distance_strategy == DistanceStrategy.MAX_INNER_PRODUCT:323 index = (324 scann.scann_ops_pybind.builder(vector, 1, "dot_product")325 .score_brute_force()326 .build()327 )328 else:329 # Default to L2, currently other metric types not initialized.330 index = (331 scann.scann_ops_pybind.builder(vector, 1, "squared_l2")332 .score_brute_force()333 .build()334 )335 documents = []336 if ids is None:337 ids = [str(uuid.uuid4()) for _ in texts]338 for i, text in enumerate(texts):339 metadata = metadatas[i] if metadatas else {}340 documents.append(Document(page_content=text, metadata=metadata))341 index_to_id = dict(enumerate(ids))342 343 if len(index_to_id) != len(documents):344 raise Exception(345 f"{len(index_to_id)} ids provided for {len(documents)} documents."346 " Each document should have an id."347 )348 349 docstore = InMemoryDocstore(dict(zip(index_to_id.values(), documents)))350 return cls(351 embedding,352 index,353 docstore,354 index_to_id,355 normalize_L2=normalize_L2,356 **kwargs,357 )358 359 @classmethod360 def from_texts(361 cls,362 texts: List[str],363 embedding: Embeddings,364 metadatas: Optional[List[dict]] = None,365 ids: Optional[List[str]] = None,366 **kwargs: Any,367 ) -> ScaNN:368 """Construct ScaNN wrapper from raw documents.369 370 This is a user friendly interface that:371 1. Embeds documents.372 2. Creates an in memory docstore373 3. Initializes the ScaNN database374 375 This is intended to be a quick way to get started.376 377 Example:378 .. code-block:: python379 380 from langchain_community.vectorstores import ScaNN381 from langchain_community.embeddings import OpenAIEmbeddings382 embeddings = OpenAIEmbeddings()383 scann = ScaNN.from_texts(texts, embeddings)384 """385 embeddings = embedding.embed_documents(texts)386 return cls.__from(387 texts,388 embeddings,389 embedding,390 metadatas=metadatas,391 ids=ids,392 **kwargs,393 )394 395 @classmethod396 def from_embeddings(397 cls,398 text_embeddings: List[Tuple[str, List[float]]],399 embedding: Embeddings,400 metadatas: Optional[List[dict]] = None,401 ids: Optional[List[str]] = None,402 **kwargs: Any,403 ) -> ScaNN:404 """Construct ScaNN wrapper from raw documents.405 406 This is a user friendly interface that:407 1. Embeds documents.408 2. Creates an in memory docstore409 3. Initializes the ScaNN database410 411 This is intended to be a quick way to get started.412 413 Example:414 .. code-block:: python415 416 from langchain_community.vectorstores import ScaNN417 from langchain_community.embeddings import OpenAIEmbeddings418 embeddings = OpenAIEmbeddings()419 text_embeddings = embeddings.embed_documents(texts)420 text_embedding_pairs = list(zip(texts, text_embeddings))421 scann = ScaNN.from_embeddings(text_embedding_pairs, embeddings)422 """423 texts = [t[0] for t in text_embeddings]424 embeddings = [t[1] for t in text_embeddings]425 return cls.__from(426 texts,427 embeddings,428 embedding,429 metadatas=metadatas,430 ids=ids,431 **kwargs,432 )433 434 def save_local(self, folder_path: str, index_name: str = "index") -> None:435 """Save ScaNN index, docstore, and index_to_docstore_id to disk.436 437 Args:438 folder_path: folder path to save index, docstore,439 and index_to_docstore_id to.440 """441 path = Path(folder_path)442 scann_path = path / "{index_name}.scann".format(index_name=index_name)443 scann_path.mkdir(exist_ok=True, parents=True)444 445 # save index separately since it is not picklable446 self.index.serialize(str(scann_path))447 448 # save docstore and index_to_docstore_id449 with open(path / "{index_name}.pkl".format(index_name=index_name), "wb") as f:450 pickle.dump((self.docstore, self.index_to_docstore_id), f)451 452 @classmethod453 def load_local(454 cls,455 folder_path: str,456 embedding: Embeddings,457 index_name: str = "index",458 *,459 allow_dangerous_deserialization: bool = False,460 **kwargs: Any,461 ) -> ScaNN:462 """Load ScaNN index, docstore, and index_to_docstore_id from disk.463 464 Args:465 folder_path: folder path to load index, docstore,466 and index_to_docstore_id from.467 embedding: Embeddings to use when generating queries468 index_name: for saving with a specific index file name469 allow_dangerous_deserialization: whether to allow deserialization470 of the data which involves loading a pickle file.471 Pickle files can be modified by malicious actors to deliver a472 malicious payload that results in execution of473 arbitrary code on your machine.474 """475 if not allow_dangerous_deserialization:476 raise ValueError(477 "The de-serialization relies loading a pickle file. "478 "Pickle files can be modified to deliver a malicious payload that "479 "results in execution of arbitrary code on your machine."480 "You will need to set `allow_dangerous_deserialization` to `True` to "481 "enable deserialization. If you do this, make sure that you "482 "trust the source of the data. For example, if you are loading a "483 "file that you created, and know that no one else has modified the "484 "file, then this is safe to do. Do not set this to `True` if you are "485 "loading a file from an untrusted source (e.g., some random site on "486 "the internet.)."487 )488 path = Path(folder_path)489 scann_path = path / "{index_name}.scann".format(index_name=index_name)490 scann_path.mkdir(exist_ok=True, parents=True)491 # load index separately since it is not picklable492 scann = guard_import("scann")493 index = scann.scann_ops_pybind.load_searcher(str(scann_path))494 495 # load docstore and index_to_docstore_id496 with open(path / "{index_name}.pkl".format(index_name=index_name), "rb") as f:497 (498 docstore,499 index_to_docstore_id,500 ) = pickle.load( # ignore[pickle]: explicit-opt-in501 f502 )503 504 return cls(embedding, index, docstore, index_to_docstore_id, **kwargs)505 506 def _select_relevance_score_fn(self) -> Callable[[float], float]:507 """508 The 'correct' relevance function509 may differ depending on a few things, including:510 - the distance / similarity metric used by the VectorStore511 - the scale of your embeddings (OpenAI's are unit normed. Many others are not!)512 - embedding dimensionality513 - etc.514 """515 if self.override_relevance_score_fn is not None:516 return self.override_relevance_score_fn517 518 # Default strategy is to rely on distance strategy provided in519 # vectorstore constructor520 if self.distance_strategy == DistanceStrategy.MAX_INNER_PRODUCT:521 return self._max_inner_product_relevance_score_fn522 elif self.distance_strategy == DistanceStrategy.EUCLIDEAN_DISTANCE:523 # Default behavior is to use euclidean distance relevancy524 return self._euclidean_relevance_score_fn525 else:526 raise ValueError(527 "Unknown distance strategy, must be cosine, max_inner_product,"528 " or euclidean"529 )530 531 def _similarity_search_with_relevance_scores(532 self,533 query: str,534 k: int = 4,535 filter: Optional[Dict[str, Any]] = None,536 fetch_k: int = 20,537 **kwargs: Any,538 ) -> List[Tuple[Document, float]]:539 """Return docs and their similarity scores on a scale from 0 to 1."""540 # Pop score threshold so that only relevancy scores, not raw scores, are541 # filtered.542 score_threshold = kwargs.pop("score_threshold", None)543 relevance_score_fn = self._select_relevance_score_fn()544 if relevance_score_fn is None:545 raise ValueError(546 "normalize_score_fn must be provided to"547 " ScaNN constructor to normalize scores"548 )549 docs_and_scores = self.similarity_search_with_score(550 query,551 k=k,552 filter=filter,553 fetch_k=fetch_k,554 **kwargs,555 )556 docs_and_rel_scores = [557 (doc, relevance_score_fn(score)) for doc, score in docs_and_scores558 ]559 if score_threshold is not None:560 docs_and_rel_scores = [561 (doc, similarity)562 for doc, similarity in docs_and_rel_scores563 if similarity >= score_threshold564 ]565 return docs_and_rel_scores566 