codekingpro/portable-devtools
114k
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 