codekingpro/portable-devtools
114k
1import importlib2import importlib.metadata3from typing import Any, Dict, List, Literal, Optional, Sequence, cast4 5import numpy as np6from langchain_core.embeddings import Embeddings7from langchain_core.utils import pre_init8from pydantic import BaseModel, ConfigDict9 10MIN_VERSION = "0.2.0"11 12 13class FastEmbedEmbeddings(BaseModel, Embeddings):14 """Qdrant FastEmbedding models.15 16 FastEmbed is a lightweight, fast, Python library built for embedding generation.17 See more documentation at:18 * https://github.com/qdrant/fastembed/19 * https://qdrant.github.io/fastembed/20 21 To use this class, you must install the `fastembed` Python package.22 23 `pip install fastembed`24 Example:25 from langchain_community.embeddings import FastEmbedEmbeddings26 fastembed = FastEmbedEmbeddings()27 """28 29 model_name: str = "BAAI/bge-small-en-v1.5"30 """Name of the FastEmbedding model to use31 Defaults to "BAAI/bge-small-en-v1.5"32 Find the list of supported models at33 https://qdrant.github.io/fastembed/examples/Supported_Models/34 """35 36 max_length: int = 51237 """The maximum number of tokens. Defaults to 512.38 Unknown behavior for values > 512.39 """40 41 cache_dir: Optional[str] = None42 """The path to the cache directory.43 Defaults to `local_cache` in the parent directory44 """45 46 threads: Optional[int] = None47 """The number of threads single onnxruntime session can use.48 Defaults to None49 """50 51 doc_embed_type: Literal["default", "passage"] = "default"52 """Type of embedding to use for documents53 The available options are: "default" and "passage"54 """55 56 batch_size: int = 25657 """Batch size for encoding. Higher values will use more memory, but be faster.58 Defaults to 256.59 """60 61 parallel: Optional[int] = None62 """If `>1`, parallel encoding is used, recommended for encoding of large datasets.63 If `0`, use all available cores.64 If `None`, don't use data-parallel processing, use default onnxruntime threading.65 Defaults to `None`.66 """67 68 providers: Optional[Sequence[Any]] = None69 """List of ONNX execution providers. Use `["CUDAExecutionProvider"]` to enable the70 use of GPU when generating embeddings. This requires to install `fastembed-gpu`71 instead of `fastembed`. See https://qdrant.github.io/fastembed/examples/FastEmbed_GPU72 for more details.73 Defaults to `None`.74 """75 76 model: Any = None # : :meta private:77 78 model_config = ConfigDict(extra="allow", protected_namespaces=())79 80 @pre_init81 def validate_environment(cls, values: Dict) -> Dict:82 """Validate that FastEmbed has been installed."""83 model_name = values.get("model_name")84 max_length = values.get("max_length")85 cache_dir = values.get("cache_dir")86 threads = values.get("threads")87 providers = values.get("providers")88 pkg_to_install = (89 "fastembed-gpu"90 if providers and "CUDAExecutionProvider" in providers91 else "fastembed"92 )93 94 try:95 fastembed = importlib.import_module("fastembed")96 97 except ModuleNotFoundError:98 raise ImportError(99 "Could not import 'fastembed' Python package. "100 f"Please install it with `pip install {pkg_to_install}`."101 )102 103 if importlib.metadata.version(pkg_to_install) < MIN_VERSION:104 raise ImportError(105 f"FastEmbedEmbeddings requires "106 f'`pip install -U "{pkg_to_install}>={MIN_VERSION}"`.'107 )108 109 values["model"] = fastembed.TextEmbedding(110 model_name=model_name,111 max_length=max_length,112 cache_dir=cache_dir,113 threads=threads,114 providers=providers,115 )116 return values117 118 def embed_documents(self, texts: List[str]) -> List[List[float]]:119 """Generate embeddings for documents using FastEmbed.120 121 Args:122 texts: The list of texts to embed.123 124 Returns:125 List of embeddings, one for each text.126 """127 embeddings: List[np.ndarray]128 if self.doc_embed_type == "passage":129 embeddings = self.model.passage_embed(130 texts, batch_size=self.batch_size, parallel=self.parallel131 )132 else:133 embeddings = self.model.embed(134 texts, batch_size=self.batch_size, parallel=self.parallel135 )136 return [cast(List[float], e.tolist()) for e in embeddings]137 138 def embed_query(self, text: str) -> List[float]:139 """Generate query embeddings using FastEmbed.140 141 Args:142 text: The text to embed.143 144 Returns:145 Embeddings for the text.146 """147 query_embeddings: np.ndarray = next(148 self.model.query_embed(149 text, batch_size=self.batch_size, parallel=self.parallel150 )151 )152 return cast(List[float], query_embeddings.tolist())153 