codekingpro/portable-devtools
114k
1from typing import Any, Dict, List, Optional2 3from langchain_core.embeddings import Embeddings4from langchain_core.utils import get_from_dict_or_env5from packaging.version import parse6from pydantic import BaseModel, ConfigDict, model_validator7from typing_extensions import Self8 9__all__ = ["GradientEmbeddings"]10 11 12class GradientEmbeddings(BaseModel, Embeddings):13 """Gradient.ai Embedding models.14 15 GradientLLM is a class to interact with Embedding Models on gradient.ai16 17 To use, set the environment variable ``GRADIENT_ACCESS_TOKEN`` with your18 API token and ``GRADIENT_WORKSPACE_ID`` for your gradient workspace,19 or alternatively provide them as keywords to the constructor of this class.20 21 Example:22 .. code-block:: python23 24 from langchain_community.embeddings import GradientEmbeddings25 GradientEmbeddings(26 model="bge-large",27 gradient_workspace_id="12345614fc0_workspace",28 gradient_access_token="gradientai-access_token",29 )30 """31 32 model: str33 "Underlying gradient.ai model id."34 35 gradient_workspace_id: Optional[str] = None36 "Underlying gradient.ai workspace_id."37 38 gradient_access_token: Optional[str] = None39 """gradient.ai API Token, which can be generated by going to40 https://auth.gradient.ai/select-workspace41 and selecting "Access tokens" under the profile drop-down.42 """43 44 gradient_api_url: str = "https://api.gradient.ai/api"45 """Endpoint URL to use."""46 47 query_prompt_for_retrieval: Optional[str] = None48 """Query pre-prompt"""49 50 client: Any = None #: :meta private:51 """Gradient client."""52 53 # LLM call kwargs54 model_config = ConfigDict(55 extra="forbid",56 )57 58 @model_validator(mode="before")59 @classmethod60 def validate_environment(cls, values: Dict) -> Any:61 """Validate that api key and python package exists in environment."""62 63 values["gradient_access_token"] = get_from_dict_or_env(64 values, "gradient_access_token", "GRADIENT_ACCESS_TOKEN"65 )66 values["gradient_workspace_id"] = get_from_dict_or_env(67 values, "gradient_workspace_id", "GRADIENT_WORKSPACE_ID"68 )69 70 values["gradient_api_url"] = get_from_dict_or_env(71 values,72 "gradient_api_url",73 "GRADIENT_API_URL",74 default="https://api.gradient.ai/api",75 )76 return values77 78 @model_validator(mode="after")79 def post_init(self) -> Self:80 try:81 import gradientai82 except ImportError:83 raise ImportError(84 'GradientEmbeddings requires `pip install -U "gradientai>=1.4.0"`.'85 )86 87 if parse(gradientai.__version__) < parse("1.4.0"):88 raise ImportError(89 'GradientEmbeddings requires `pip install -U "gradientai>=1.4.0"`.'90 )91 92 gradient = gradientai.Gradient(93 access_token=self.gradient_access_token,94 workspace_id=self.gradient_workspace_id,95 host=self.gradient_api_url,96 )97 self.client = gradient.get_embeddings_model(slug=self.model)98 return self99 100 def embed_documents(self, texts: List[str]) -> List[List[float]]:101 """Call out to Gradient's embedding endpoint.102 103 Args:104 texts: The list of texts to embed.105 106 Returns:107 List of embeddings, one for each text.108 """109 inputs = [{"input": text} for text in texts]110 111 result = self.client.embed(inputs=inputs).embeddings112 113 return [e.embedding for e in result]114 115 async def aembed_documents(self, texts: List[str]) -> List[List[float]]:116 """Async call out to Gradient's embedding endpoint.117 118 Args:119 texts: The list of texts to embed.120 121 Returns:122 List of embeddings, one for each text.123 """124 inputs = [{"input": text} for text in texts]125 126 result = (await self.client.aembed(inputs=inputs)).embeddings127 128 return [e.embedding for e in result]129 130 def embed_query(self, text: str) -> List[float]:131 """Call out to Gradient's embedding endpoint.132 133 Args:134 text: The text to embed.135 136 Returns:137 Embeddings for the text.138 """139 query = (140 f"{self.query_prompt_for_retrieval} {text}"141 if self.query_prompt_for_retrieval142 else text143 )144 return self.embed_documents([query])[0]145 146 async def aembed_query(self, text: str) -> List[float]:147 """Async call out to Gradient's embedding endpoint.148 149 Args:150 text: The text to embed.151 152 Returns:153 Embeddings for the text.154 """155 query = (156 f"{self.query_prompt_for_retrieval} {text}"157 if self.query_prompt_for_retrieval158 else text159 )160 embeddings = await self.aembed_documents([query])161 return embeddings[0]162 163 164class TinyAsyncGradientEmbeddingClient: #: :meta private:165 """Deprecated, TinyAsyncGradientEmbeddingClient was removed.166 167 This class is just for backwards compatibility with older versions168 of langchain_community.169 It might be entirely removed in the future.170 """171 172 def __init__(self, *args: Any, **kwargs: Any) -> None:173 raise ValueError("Deprecated,TinyAsyncGradientEmbeddingClient was removed.")174 