codekingpro/portable-devtools
114k
1import asyncio2import json3import os4from typing import Any, Dict, List, Optional5 6import numpy as np7from langchain_core._api.deprecation import deprecated8from langchain_core.embeddings import Embeddings9from langchain_core.runnables.config import run_in_executor10from pydantic import BaseModel, ConfigDict, model_validator11from typing_extensions import Self12 13 14@deprecated(15 since="0.2.11",16 removal="1.0",17 alternative_import="langchain_aws.BedrockEmbeddings",18)19class BedrockEmbeddings(BaseModel, Embeddings):20 """Bedrock embedding models.21 22 To authenticate, the AWS client uses the following methods to23 automatically load credentials:24 https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html25 26 If a specific credential profile should be used, you must pass27 the name of the profile from the ~/.aws/credentials file that is to be used.28 29 Make sure the credentials / roles used have the required policies to30 access the Bedrock service.31 """32 33 """34 Example:35 .. code-block:: python36 37 from langchain_community.bedrock_embeddings import BedrockEmbeddings38 39 region_name ="us-east-1"40 credentials_profile_name = "default"41 model_id = "amazon.titan-embed-text-v1"42 43 be = BedrockEmbeddings(44 credentials_profile_name=credentials_profile_name,45 region_name=region_name,46 model_id=model_id47 )48 """49 50 client: Any = None #: :meta private:51 """Bedrock client."""52 region_name: Optional[str] = None53 """The aws region e.g., `us-west-2`. Fallsback to AWS_DEFAULT_REGION env variable54 or region specified in ~/.aws/config in case it is not provided here.55 """56 57 credentials_profile_name: Optional[str] = None58 """The name of the profile in the ~/.aws/credentials or ~/.aws/config files, which59 has either access keys or role information specified.60 If not specified, the default credential profile or, if on an EC2 instance,61 credentials from IMDS will be used.62 See: https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html63 """64 65 model_id: str = "amazon.titan-embed-text-v1"66 """Id of the model to call, e.g., amazon.titan-embed-text-v1, this is67 equivalent to the modelId property in the list-foundation-models api"""68 69 model_kwargs: Optional[Dict] = None70 """Keyword arguments to pass to the model."""71 72 endpoint_url: Optional[str] = None73 """Needed if you don't want to default to us-east-1 endpoint"""74 75 normalize: bool = False76 """Whether the embeddings should be normalized to unit vectors"""77 78 model_config = ConfigDict(extra="forbid", protected_namespaces=())79 80 @model_validator(mode="after")81 def validate_environment(self) -> Self:82 """Validate that AWS credentials to and python package exists in environment."""83 84 if self.client is not None:85 return self86 87 try:88 import boto389 90 if self.credentials_profile_name is not None:91 session = boto3.Session(profile_name=self.credentials_profile_name)92 else:93 # use default credentials94 session = boto3.Session()95 96 client_params = {}97 if self.region_name:98 client_params["region_name"] = self.region_name99 100 if self.endpoint_url:101 client_params["endpoint_url"] = self.endpoint_url102 103 self.client = session.client("bedrock-runtime", **client_params)104 105 except ImportError:106 raise ImportError(107 "Could not import boto3 python package. "108 "Please install it with `pip install boto3`."109 )110 except Exception as e:111 raise ValueError(112 "Could not load credentials to authenticate with AWS client. "113 "Please check that credentials in the specified "114 f"profile name are valid. Bedrock error: {e}"115 ) from e116 117 return self118 119 def _embedding_func(self, text: str) -> List[float]:120 """Call out to Bedrock embedding endpoint."""121 # replace newlines, which can negatively affect performance.122 text = text.replace(os.linesep, " ")123 124 # format input body for provider125 provider = self.model_id.split(".")[0]126 _model_kwargs = self.model_kwargs or {}127 input_body = {**_model_kwargs}128 if provider == "cohere":129 if "input_type" not in input_body.keys():130 input_body["input_type"] = "search_document"131 input_body["texts"] = [text]132 else:133 # includes common provider == "amazon"134 input_body["inputText"] = text135 body = json.dumps(input_body)136 137 try:138 # invoke bedrock API139 response = self.client.invoke_model(140 body=body,141 modelId=self.model_id,142 accept="application/json",143 contentType="application/json",144 )145 146 # format output based on provider147 response_body = json.loads(response.get("body").read())148 if provider == "cohere":149 return response_body.get("embeddings")[0]150 else:151 # includes common provider == "amazon"152 return response_body.get("embedding")153 except Exception as e:154 raise ValueError(f"Error raised by inference endpoint: {e}")155 156 def _normalize_vector(self, embeddings: List[float]) -> List[float]:157 """Normalize the embedding to a unit vector."""158 emb = np.array(embeddings)159 norm_emb = emb / np.linalg.norm(emb)160 return norm_emb.tolist()161 162 def embed_documents(self, texts: List[str]) -> List[List[float]]:163 """Compute doc embeddings using a Bedrock model.164 165 Args:166 texts: The list of texts to embed167 168 Returns:169 List of embeddings, one for each text.170 """171 results = []172 for text in texts:173 response = self._embedding_func(text)174 175 if self.normalize:176 response = self._normalize_vector(response)177 178 results.append(response)179 180 return results181 182 def embed_query(self, text: str) -> List[float]:183 """Compute query embeddings using a Bedrock model.184 185 Args:186 text: The text to embed.187 188 Returns:189 Embeddings for the text.190 """191 embedding = self._embedding_func(text)192 193 if self.normalize:194 return self._normalize_vector(embedding)195 196 return embedding197 198 async def aembed_query(self, text: str) -> List[float]:199 """Asynchronous compute query embeddings using a Bedrock model.200 201 Args:202 text: The text to embed.203 204 Returns:205 Embeddings for the text.206 """207 208 return await run_in_executor(None, self.embed_query, text)209 210 async def aembed_documents(self, texts: List[str]) -> List[List[float]]:211 """Asynchronous compute doc embeddings using a Bedrock model.212 213 Args:214 texts: The list of texts to embed215 216 Returns:217 List of embeddings, one for each text.218 """219 220 result = await asyncio.gather(*[self.aembed_query(text) for text in texts])221 222 return list(result)223 