codekingpro/portable-devtools
114k
1from __future__ import annotations2 3from typing import Any, Dict, Iterator, List4from urllib.parse import urlparse5 6from langchain_core.embeddings import Embeddings7from pydantic import BaseModel, PrivateAttr8 9 10def _chunk(texts: List[str], size: int) -> Iterator[List[str]]:11 for i in range(0, len(texts), size):12 yield texts[i : i + size]13 14 15class MlflowEmbeddings(Embeddings, BaseModel):16 """Embedding LLMs in MLflow.17 18 To use, you should have the `mlflow[genai]` python package installed.19 For more information, see https://mlflow.org/docs/latest/llms/deployments.20 21 Example:22 .. code-block:: python23 24 from langchain_community.embeddings import MlflowEmbeddings25 26 embeddings = MlflowEmbeddings(27 target_uri="http://localhost:5000",28 endpoint="embeddings",29 )30 """31 32 endpoint: str33 """The endpoint to use."""34 target_uri: str35 """The target URI to use."""36 _client: Any = PrivateAttr()37 """The parameters to use for queries."""38 query_params: Dict[str, str] = {}39 """The parameters to use for documents."""40 documents_params: Dict[str, str] = {}41 42 def __init__(self, **kwargs: Any):43 super().__init__(**kwargs)44 self._validate_uri()45 try:46 from mlflow.deployments import get_deploy_client47 48 self._client = get_deploy_client(self.target_uri)49 except ImportError as e:50 raise ImportError(51 "Failed to create the client. "52 f"Please run `pip install mlflow{self._mlflow_extras}` to install "53 "required dependencies."54 ) from e55 56 @property57 def _mlflow_extras(self) -> str:58 return "[genai]"59 60 def _validate_uri(self) -> None:61 if self.target_uri == "databricks":62 return63 allowed = ["http", "https", "databricks"]64 if urlparse(self.target_uri).scheme not in allowed:65 raise ValueError(66 f"Invalid target URI: {self.target_uri}. "67 f"The scheme must be one of {allowed}."68 )69 70 def embed(self, texts: List[str], params: Dict[str, str]) -> List[List[float]]:71 embeddings: List[List[float]] = []72 for txt in _chunk(texts, 20):73 resp = self._client.predict(74 endpoint=self.endpoint,75 inputs={"input": txt, **params},76 )77 embeddings.extend(r["embedding"] for r in resp["data"])78 return embeddings79 80 def embed_documents(self, texts: List[str]) -> List[List[float]]:81 return self.embed(texts, params=self.documents_params)82 83 def embed_query(self, text: str) -> List[float]:84 return self.embed([text], params=self.query_params)[0]85 86 87class MlflowCohereEmbeddings(MlflowEmbeddings):88 """Cohere embedding LLMs in MLflow."""89 90 query_params: Dict[str, str] = {"input_type": "search_query"}91 documents_params: Dict[str, str] = {"input_type": "search_document"}92 