Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
gradient_ai.py174 linesDownload Raw Back to embeddings
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 
codekingpro/portable-devtools · Team Ai