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