codekingpro/portable-devtools
114k
1"""Wrapper around the Baidu vector database."""2 3from __future__ import annotations4 5import json6import logging7import time8from typing import Any, Dict, Iterable, List, Optional, Tuple9 10import numpy as np11from langchain_core.documents import Document12from langchain_core.embeddings import Embeddings13from langchain_core.utils import guard_import14from langchain_core.vectorstores import VectorStore15 16from langchain_community.vectorstores.utils import maximal_marginal_relevance17 18logger = logging.getLogger(__name__)19 20 21class ConnectionParams:22 """Baidu VectorDB Connection params.23 24 See the following documentation for details:25 https://cloud.baidu.com/doc/VDB/s/6lrsob0wy26 27 Attribute:28 endpoint (str) : The access address of the vector database server29 that the client needs to connect to.30 api_key (str): API key for client to access the vector database server,31 which is used for authentication.32 account (str) : Account for client to access the vector database server.33 connection_timeout_in_mills (int) : Request Timeout.34 """35 36 def __init__(37 self,38 endpoint: str,39 api_key: str,40 account: str = "root",41 connection_timeout_in_mills: int = 50 * 1000,42 ):43 self.endpoint = endpoint44 self.api_key = api_key45 self.account = account46 self.connection_timeout_in_mills = connection_timeout_in_mills47 48 49class TableParams:50 """Baidu VectorDB table params.51 52 See the following documentation for details:53 https://cloud.baidu.com/doc/VDB/s/mlrsob0p654 """55 56 def __init__(57 self,58 dimension: int,59 replication: int = 3,60 partition: int = 1,61 index_type: str = "HNSW",62 metric_type: str = "L2",63 params: Optional[Dict] = None,64 ):65 self.dimension = dimension66 self.replication = replication67 self.partition = partition68 self.index_type = index_type69 self.metric_type = metric_type70 self.params = params71 72 73class BaiduVectorDB(VectorStore):74 """Baidu VectorDB as a vector store.75 76 In order to use this you need to have a database instance.77 See the following documentation for details:78 https://cloud.baidu.com/doc/VDB/index.html79 """80 81 field_id: str = "id"82 field_vector: str = "vector"83 field_text: str = "text"84 field_metadata: str = "metadata"85 86 index_vector: str = "vector_idx"87 88 def __init__(89 self,90 embedding: Embeddings,91 connection_params: ConnectionParams,92 table_params: TableParams = TableParams(128),93 database_name: str = "LangChainDatabase",94 table_name: str = "LangChainTable",95 drop_old: Optional[bool] = False,96 ):97 pymochow = guard_import("pymochow")98 configuration = guard_import("pymochow.configuration")99 auth = guard_import("pymochow.auth.bce_credentials")100 self.mochowtable = guard_import("pymochow.model.table")101 self.mochowenum = guard_import("pymochow.model.enum")102 self.embedding_func = embedding103 self.table_params = table_params104 config = configuration.Configuration(105 credentials=auth.BceCredentials(106 connection_params.account, connection_params.api_key107 ),108 endpoint=connection_params.endpoint,109 connection_timeout_in_mills=connection_params.connection_timeout_in_mills,110 )111 self.vdb_client = pymochow.MochowClient(config)112 db_list = self.vdb_client.list_databases()113 db_exist: bool = False114 for db in db_list:115 if database_name == db.database_name:116 db_exist = True117 break118 if db_exist:119 self.database = self.vdb_client.database(database_name)120 else:121 self.database = self.vdb_client.create_database(database_name)122 try:123 self.table = self.database.describe_table(table_name)124 if drop_old:125 self.database.drop_table(table_name)126 self._create_table(table_name)127 except pymochow.exception.ServerError:128 self._create_table(table_name)129 130 def _create_table(self, table_name: str) -> None:131 schema = guard_import("pymochow.model.schema")132 index_type = None133 for k, v in self.mochowenum.IndexType.__members__.items():134 if k == self.table_params.index_type:135 index_type = v136 if index_type is None:137 raise ValueError("unsupported index_type")138 metric_type = None139 for k, v in self.mochowenum.MetricType.__members__.items():140 if k == self.table_params.metric_type:141 metric_type = v142 if metric_type is None:143 raise ValueError("unsupported metric_type")144 if self.table_params.params is None:145 params = schema.HNSWParams(m=16, efconstruction=200)146 else:147 params = schema.HNSWParams(148 m=self.table_params.params.get("M", 16),149 efconstruction=self.table_params.params.get("efConstruction", 200),150 )151 fields = []152 fields.append(153 schema.Field(154 self.field_id,155 self.mochowenum.FieldType.STRING,156 primary_key=True,157 partition_key=True,158 auto_increment=False,159 not_null=True,160 )161 )162 fields.append(163 schema.Field(164 self.field_vector,165 self.mochowenum.FieldType.FLOAT_VECTOR,166 dimension=self.table_params.dimension,167 not_null=True,168 )169 )170 fields.append(schema.Field(self.field_text, self.mochowenum.FieldType.STRING))171 fields.append(172 schema.Field(self.field_metadata, self.mochowenum.FieldType.STRING)173 )174 indexes = []175 indexes.append(176 schema.VectorIndex(177 index_name=self.index_vector,178 index_type=index_type,179 field=self.field_vector,180 metric_type=metric_type,181 params=params,182 )183 )184 185 self.table = self.database.create_table(186 table_name=table_name,187 replication=self.table_params.replication,188 partition=self.mochowtable.Partition(189 partition_num=self.table_params.partition190 ),191 schema=schema.Schema(fields=fields, indexes=indexes),192 )193 194 while True:195 time.sleep(1)196 table = self.database.describe_table(table_name)197 if table.state == self.mochowenum.TableState.NORMAL:198 break199 200 @property201 def embeddings(self) -> Embeddings:202 return self.embedding_func203 204 @classmethod205 def from_texts(206 cls,207 texts: List[str],208 embedding: Embeddings,209 metadatas: Optional[List[dict]] = None,210 connection_params: Optional[ConnectionParams] = None,211 table_params: Optional[TableParams] = None,212 database_name: str = "LangChainDatabase",213 table_name: str = "LangChainTable",214 drop_old: Optional[bool] = False,215 **kwargs: Any,216 ) -> BaiduVectorDB:217 """Create a table, indexes it with HNSW, and insert data."""218 if len(texts) == 0:219 raise ValueError("texts is empty")220 if connection_params is None:221 raise ValueError("connection_params is empty")222 try:223 embeddings = embedding.embed_documents(texts[0:1])224 except NotImplementedError:225 embeddings = [embedding.embed_query(texts[0])]226 dimension = len(embeddings[0])227 if table_params is None:228 table_params = TableParams(dimension=dimension)229 else:230 table_params.dimension = dimension231 vector_db = cls(232 embedding=embedding,233 connection_params=connection_params,234 table_params=table_params,235 database_name=database_name,236 table_name=table_name,237 drop_old=drop_old,238 )239 vector_db.add_texts(texts=texts, metadatas=metadatas)240 return vector_db241 242 def add_texts(243 self,244 texts: Iterable[str],245 metadatas: Optional[List[dict]] = None,246 batch_size: int = 1000,247 **kwargs: Any,248 ) -> List[str]:249 """Insert text data into Baidu VectorDB."""250 texts = list(texts)251 try:252 embeddings = self.embedding_func.embed_documents(texts)253 except NotImplementedError:254 embeddings = [self.embedding_func.embed_query(x) for x in texts]255 if len(embeddings) == 0:256 logger.debug("Nothing to insert, skipping.")257 return []258 pks: list[str] = []259 total_count = len(embeddings)260 for start in range(0, total_count, batch_size):261 # Grab end index262 rows = []263 end = min(start + batch_size, total_count)264 for id in range(start, end, 1):265 metadata = "{}"266 if metadatas is not None:267 metadata = json.dumps(metadatas[id])268 row = self.mochowtable.Row(269 id="{}-{}-{}".format(time.time_ns(), hash(texts[id]), id),270 vector=[float(num) for num in embeddings[id]],271 text=texts[id],272 metadata=metadata,273 )274 rows.append(row)275 pks.append(str(id))276 self.table.upsert(rows=rows)277 # need rebuild vindex after upsert278 self.table.rebuild_index(self.index_vector)279 while True:280 time.sleep(2)281 index = self.table.describe_index(self.index_vector)282 if index.state == self.mochowenum.IndexState.NORMAL:283 break284 return pks285 286 def similarity_search(287 self,288 query: str,289 k: int = 4,290 param: Optional[dict] = None,291 expr: Optional[str] = None,292 **kwargs: Any,293 ) -> List[Document]:294 """Perform a similarity search against the query string."""295 res = self.similarity_search_with_score(296 query=query, k=k, param=param, expr=expr, **kwargs297 )298 return [doc for doc, _ in res]299 300 def similarity_search_with_score(301 self,302 query: str,303 k: int = 4,304 param: Optional[dict] = None,305 expr: Optional[str] = None,306 **kwargs: Any,307 ) -> List[Tuple[Document, float]]:308 """Perform a search on a query string and return results with score."""309 # Embed the query text.310 embedding = self.embedding_func.embed_query(query)311 res = self._similarity_search_with_score(312 embedding=embedding, k=k, param=param, expr=expr, **kwargs313 )314 return res315 316 def similarity_search_by_vector(317 self,318 embedding: List[float],319 k: int = 4,320 param: Optional[dict] = None,321 expr: Optional[str] = None,322 **kwargs: Any,323 ) -> List[Document]:324 """Perform a similarity search against the query string."""325 res = self._similarity_search_with_score(326 embedding=embedding, k=k, param=param, expr=expr, **kwargs327 )328 return [doc for doc, _ in res]329 330 def _similarity_search_with_score(331 self,332 embedding: List[float],333 k: int = 4,334 param: Optional[dict] = None,335 expr: Optional[str] = None,336 **kwargs: Any,337 ) -> List[Tuple[Document, float]]:338 """Perform a search on a query string and return results with score."""339 ef = 10 if param is None else param.get("ef", 10)340 341 anns = self.mochowtable.AnnSearch(342 vector_field=self.field_vector,343 vector_floats=[float(num) for num in embedding],344 params=self.mochowtable.HNSWSearchParams(ef=ef, limit=k),345 filter=expr,346 )347 res = self.table.search(anns=anns)348 349 rows = [[item] for item in res.rows]350 # Organize results.351 ret: List[Tuple[Document, float]] = []352 if rows is None or len(rows) == 0:353 return ret354 for row in rows:355 for result in row:356 row_data = result.get("row", {})357 meta = row_data.get(self.field_metadata)358 if meta is not None:359 meta = json.loads(meta)360 doc = Document(361 page_content=row_data.get(self.field_text), metadata=meta362 )363 pair = (doc, result.get("score", 0.0))364 ret.append(pair)365 return ret366 367 def max_marginal_relevance_search(368 self,369 query: str,370 k: int = 4,371 fetch_k: int = 20,372 lambda_mult: float = 0.5,373 param: Optional[dict] = None,374 expr: Optional[str] = None,375 **kwargs: Any,376 ) -> List[Document]:377 """Perform a search and return results that are reordered by MMR."""378 embedding = self.embedding_func.embed_query(query)379 return self._max_marginal_relevance_search(380 embedding=embedding,381 k=k,382 fetch_k=fetch_k,383 lambda_mult=lambda_mult,384 param=param,385 expr=expr,386 **kwargs,387 )388 389 def _max_marginal_relevance_search(390 self,391 embedding: list[float],392 k: int = 4,393 fetch_k: int = 20,394 lambda_mult: float = 0.5,395 param: Optional[dict] = None,396 expr: Optional[str] = None,397 **kwargs: Any,398 ) -> List[Document]:399 """Perform a search and return results that are reordered by MMR."""400 ef = 10 if param is None else param.get("ef", 10)401 anns = self.mochowtable.AnnSearch(402 vector_field=self.field_vector,403 vector_floats=[float(num) for num in embedding],404 params=self.mochowtable.HNSWSearchParams(ef=ef, limit=k),405 filter=expr,406 )407 res = self.table.search(anns=anns, retrieve_vector=True)408 409 # Organize results.410 documents: List[Document] = []411 ordered_result_embeddings = []412 rows = [[item] for item in res.rows]413 if rows is None or len(rows) == 0:414 return documents415 for row in rows:416 for result in row:417 row_data = result.get("row", {})418 meta = row_data.get(self.field_metadata)419 if meta is not None:420 meta = json.loads(meta)421 doc = Document(422 page_content=row_data.get(self.field_text), metadata=meta423 )424 documents.append(doc)425 ordered_result_embeddings.append(row_data.get(self.field_vector))426 # Get the new order of results.427 new_ordering = maximal_marginal_relevance(428 np.array(embedding), ordered_result_embeddings, k=k, lambda_mult=lambda_mult429 )430 # Reorder the values and return.431 ret = []432 for x in new_ordering:433 # Function can return -1 index434 if x == -1:435 break436 else:437 ret.append(documents[x])438 return ret439 