Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
sagemaker_endpoint.py211 linesDownload Raw Back to embeddings
1from typing import Any, Dict, List, Optional2 3from langchain_core.embeddings import Embeddings4from langchain_core.utils import pre_init5from pydantic import BaseModel, ConfigDict6 7from langchain_community.llms.sagemaker_endpoint import ContentHandlerBase8 9 10class EmbeddingsContentHandler(ContentHandlerBase[List[str], List[List[float]]]):11    """Content handler for LLM class."""12 13 14class SagemakerEndpointEmbeddings(BaseModel, Embeddings):15    """Custom Sagemaker Inference Endpoints.16 17    To use, you must supply the endpoint name from your deployed18    Sagemaker model & the region where it is deployed.19 20    To authenticate, the AWS client uses the following methods to21    automatically load credentials:22    https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html23 24    If a specific credential profile should be used, you must pass25    the name of the profile from the ~/.aws/credentials file that is to be used.26 27    Make sure the credentials / roles used have the required policies to28    access the Sagemaker endpoint.29    See: https://docs.aws.amazon.com/IAM/latest/UserGuide/access_policies.html30    """31 32    """33    Example:34        .. code-block:: python35 36            from langchain_community.embeddings import SagemakerEndpointEmbeddings37            endpoint_name = (38                "my-endpoint-name"39            )40            region_name = (41                "us-west-2"42            )43            credentials_profile_name = (44                "default"45            )46            se = SagemakerEndpointEmbeddings(47                endpoint_name=endpoint_name,48                region_name=region_name,49                credentials_profile_name=credentials_profile_name50            )51 52            #Use with boto3 client53            client = boto3.client(54                        "sagemaker-runtime",55                        region_name=region_name56                    )57            se = SagemakerEndpointEmbeddings(58                endpoint_name=endpoint_name,59                client=client60            )61    """62    client: Any = None63 64    endpoint_name: str = ""65    """The name of the endpoint from the deployed Sagemaker model.66    Must be unique within an AWS Region."""67 68    region_name: str = ""69    """The aws region where the Sagemaker model is deployed, eg. `us-west-2`."""70 71    credentials_profile_name: Optional[str] = None72    """The name of the profile in the ~/.aws/credentials or ~/.aws/config files, which73    has either access keys or role information specified.74    If not specified, the default credential profile or, if on an EC2 instance,75    credentials from IMDS will be used.76    See: https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html77    """78 79    content_handler: EmbeddingsContentHandler80    """The content handler class that provides an input and81    output transform functions to handle formats between LLM82    and the endpoint.83    """84 85    """86     Example:87        .. code-block:: python88 89        from langchain_community.embeddings.sagemaker_endpoint import EmbeddingsContentHandler90 91        class ContentHandler(EmbeddingsContentHandler):92                content_type = "application/json"93                accepts = "application/json"94 95                def transform_input(self, prompts: List[str], model_kwargs: Dict) -> bytes:96                    input_str = json.dumps({prompts: prompts, **model_kwargs})97                    return input_str.encode('utf-8')98 99                def transform_output(self, output: bytes) -> List[List[float]]:100                    response_json = json.loads(output.read().decode("utf-8"))101                    return response_json["vectors"]102    """  # noqa: E501103 104    model_kwargs: Optional[Dict] = None105    """Keyword arguments to pass to the model."""106 107    endpoint_kwargs: Optional[Dict] = None108    """Optional attributes passed to the invoke_endpoint109    function. See `boto3`_. docs for more info.110    .. _boto3: <https://boto3.amazonaws.com/v1/documentation/api/latest/index.html>111    """112 113    model_config = ConfigDict(114        arbitrary_types_allowed=True, extra="forbid", protected_namespaces=()115    )116 117    @pre_init118    def validate_environment(cls, values: Dict) -> Dict:119        """Dont do anything if client provided externally"""120        if values.get("client") is not None:121            return values122 123        """Validate that AWS credentials to and python package exists in environment."""124        try:125            import boto3126 127            try:128                if values["credentials_profile_name"] is not None:129                    session = boto3.Session(130                        profile_name=values["credentials_profile_name"]131                    )132                else:133                    # use default credentials134                    session = boto3.Session()135 136                values["client"] = session.client(137                    "sagemaker-runtime", region_name=values["region_name"]138                )139 140            except Exception as e:141                raise ValueError(142                    "Could not load credentials to authenticate with AWS client. "143                    "Please check that credentials in the specified "144                    f"profile name are valid. {e}"145                ) from e146 147        except ImportError:148            raise ImportError(149                "Could not import boto3 python package. "150                "Please install it with `pip install boto3`."151            )152        return values153 154    def _embedding_func(self, texts: List[str]) -> List[List[float]]:155        """Call out to SageMaker Inference embedding endpoint."""156        # replace newlines, which can negatively affect performance.157        texts = list(map(lambda x: x.replace("\n", " "), texts))158        _model_kwargs = self.model_kwargs or {}159        _endpoint_kwargs = self.endpoint_kwargs or {}160 161        body = self.content_handler.transform_input(texts, _model_kwargs)162        content_type = self.content_handler.content_type163        accepts = self.content_handler.accepts164 165        # send request166        try:167            response = self.client.invoke_endpoint(168                EndpointName=self.endpoint_name,169                Body=body,170                ContentType=content_type,171                Accept=accepts,172                **_endpoint_kwargs,173            )174        except Exception as e:175            raise ValueError(f"Error raised by inference endpoint: {e}")176 177        return self.content_handler.transform_output(response["Body"])178 179    def embed_documents(180        self, texts: List[str], chunk_size: int = 64181    ) -> List[List[float]]:182        """Compute doc embeddings using a SageMaker Inference Endpoint.183 184        Args:185            texts: The list of texts to embed.186            chunk_size: The chunk size defines how many input texts will187                be grouped together as request. If None, will use the188                chunk size specified by the class.189 190 191        Returns:192            List of embeddings, one for each text.193        """194        results = []195        _chunk_size = len(texts) if chunk_size > len(texts) else chunk_size196        for i in range(0, len(texts), _chunk_size):197            response = self._embedding_func(texts[i : i + _chunk_size])198            results.extend(response)199        return results200 201    def embed_query(self, text: str) -> List[float]:202        """Compute query embeddings using a SageMaker inference endpoint.203 204        Args:205            text: The text to embed.206 207        Returns:208            Embeddings for the text.209        """210        return self._embedding_func([text])[0]211 
codekingpro/portable-devtools · Team Ai