codekingpro/portable-devtools
114k
1import json2from typing import Dict, Generator, List, Optional3 4import requests5from langchain_core._api.deprecation import deprecated6from langchain_core.embeddings import Embeddings7from langchain_core.utils import get_from_dict_or_env, pre_init8from pydantic import BaseModel, ConfigDict9 10 11@deprecated(12 since="0.3.16",13 removal="1.0",14 alternative_import="langchain_sambanova.SambaStudioEmbeddings",15)16class SambaStudioEmbeddings(BaseModel, Embeddings):17 """SambaNova embedding models.18 19 To use, you should have the environment variables20 ``SAMBASTUDIO_EMBEDDINGS_BASE_URL``, ``SAMBASTUDIO_EMBEDDINGS_BASE_URI``21 ``SAMBASTUDIO_EMBEDDINGS_PROJECT_ID``, ``SAMBASTUDIO_EMBEDDINGS_ENDPOINT_ID``,22 ``SAMBASTUDIO_EMBEDDINGS_API_KEY``23 set with your personal sambastudio variable or pass it as a named parameter24 to the constructor.25 26 Example:27 .. code-block:: python28 29 from langchain_community.embeddings import SambaStudioEmbeddings30 31 embeddings = SambaStudioEmbeddings(sambastudio_embeddings_base_url=base_url,32 sambastudio_embeddings_base_uri=base_uri,33 sambastudio_embeddings_project_id=project_id,34 sambastudio_embeddings_endpoint_id=endpoint_id,35 sambastudio_embeddings_api_key=api_key,36 batch_size=32)37 (or)38 39 embeddings = SambaStudioEmbeddings(batch_size=32)40 41 (or)42 43 # CoE example44 embeddings = SambaStudioEmbeddings(45 batch_size=1,46 model_kwargs={47 'select_expert':'e5-mistral-7b-instruct'48 }49 )50 """51 52 sambastudio_embeddings_base_url: str = ""53 """Base url to use"""54 55 sambastudio_embeddings_base_uri: str = ""56 """endpoint base uri"""57 58 sambastudio_embeddings_project_id: str = ""59 """Project id on sambastudio for model"""60 61 sambastudio_embeddings_endpoint_id: str = ""62 """endpoint id on sambastudio for model"""63 64 sambastudio_embeddings_api_key: str = ""65 """sambastudio api key"""66 67 model_kwargs: dict = {}68 """Key word arguments to pass to the model."""69 70 batch_size: int = 3271 """Batch size for the embedding models"""72 73 model_config = ConfigDict(protected_namespaces=())74 75 @pre_init76 def validate_environment(cls, values: Dict) -> Dict:77 """Validate that api key and python package exists in environment."""78 values["sambastudio_embeddings_base_url"] = get_from_dict_or_env(79 values, "sambastudio_embeddings_base_url", "SAMBASTUDIO_EMBEDDINGS_BASE_URL"80 )81 values["sambastudio_embeddings_base_uri"] = get_from_dict_or_env(82 values,83 "sambastudio_embeddings_base_uri",84 "SAMBASTUDIO_EMBEDDINGS_BASE_URI",85 default="api/predict/generic",86 )87 values["sambastudio_embeddings_project_id"] = get_from_dict_or_env(88 values,89 "sambastudio_embeddings_project_id",90 "SAMBASTUDIO_EMBEDDINGS_PROJECT_ID",91 )92 values["sambastudio_embeddings_endpoint_id"] = get_from_dict_or_env(93 values,94 "sambastudio_embeddings_endpoint_id",95 "SAMBASTUDIO_EMBEDDINGS_ENDPOINT_ID",96 )97 values["sambastudio_embeddings_api_key"] = get_from_dict_or_env(98 values, "sambastudio_embeddings_api_key", "SAMBASTUDIO_EMBEDDINGS_API_KEY"99 )100 return values101 102 def _get_tuning_params(self) -> str:103 """104 Get the tuning parameters to use when calling the model105 106 Returns:107 The tuning parameters as a JSON string.108 """109 if "api/v2/predict/generic" in self.sambastudio_embeddings_base_uri:110 tuning_params_dict = self.model_kwargs111 else:112 tuning_params_dict = {113 k: {"type": type(v).__name__, "value": str(v)}114 for k, v in (self.model_kwargs.items())115 }116 tuning_params = json.dumps(tuning_params_dict)117 return tuning_params118 119 def _get_full_url(self, path: str) -> str:120 """121 Return the full API URL for a given path.122 123 :param str path: the sub-path124 :returns: the full API URL for the sub-path125 :rtype: str126 """127 return f"{self.sambastudio_embeddings_base_url}/{self.sambastudio_embeddings_base_uri}/{path}" # noqa: E501128 129 def _iterate_over_batches(self, texts: List[str], batch_size: int) -> Generator:130 """Generator for creating batches in the embed documents method131 Args:132 texts (List[str]): list of strings to embed133 batch_size (int, optional): batch size to be used for the embedding model.134 Will depend on the RDU endpoint used.135 Yields:136 List[str]: list (batch) of strings of size batch size137 """138 for i in range(0, len(texts), batch_size):139 yield texts[i : i + batch_size]140 141 def embed_documents(142 self, texts: List[str], batch_size: Optional[int] = None143 ) -> List[List[float]]:144 """Returns a list of embeddings for the given sentences.145 Args:146 texts (`List[str]`): List of texts to encode147 batch_size (`int`): Batch size for the encoding148 149 Returns:150 `List[np.ndarray]` or `List[tensor]`: List of embeddings151 for the given sentences152 """153 if batch_size is None:154 batch_size = self.batch_size155 http_session = requests.Session()156 url = self._get_full_url(157 f"{self.sambastudio_embeddings_project_id}/{self.sambastudio_embeddings_endpoint_id}"158 )159 params = json.loads(self._get_tuning_params())160 embeddings = []161 162 if "api/predict/nlp" in self.sambastudio_embeddings_base_uri:163 for batch in self._iterate_over_batches(texts, batch_size):164 data = {"inputs": batch, "params": params}165 response = http_session.post(166 url,167 headers={"key": self.sambastudio_embeddings_api_key},168 json=data,169 )170 if response.status_code != 200:171 raise RuntimeError(172 f"Sambanova /complete call failed with status code "173 f"{response.status_code}.\n Details: {response.text}"174 )175 try:176 embedding = response.json()["data"]177 embeddings.extend(embedding)178 except KeyError:179 raise KeyError(180 "'data' not found in endpoint response",181 response.json(),182 )183 184 elif "api/v2/predict/generic" in self.sambastudio_embeddings_base_uri:185 for batch in self._iterate_over_batches(texts, batch_size):186 items = [187 {"id": f"item{i}", "value": item} for i, item in enumerate(batch)188 ]189 data = {"items": items, "params": params}190 response = http_session.post(191 url,192 headers={"key": self.sambastudio_embeddings_api_key},193 json=data,194 )195 if response.status_code != 200:196 raise RuntimeError(197 f"Sambanova /complete call failed with status code "198 f"{response.status_code}.\n Details: {response.text}"199 )200 try:201 embedding = [item["value"] for item in response.json()["items"]]202 embeddings.extend(embedding)203 except KeyError:204 raise KeyError(205 "'items' not found in endpoint response",206 response.json(),207 )208 209 elif "api/predict/generic" in self.sambastudio_embeddings_base_uri:210 for batch in self._iterate_over_batches(texts, batch_size):211 data = {"instances": batch, "params": params}212 response = http_session.post(213 url,214 headers={"key": self.sambastudio_embeddings_api_key},215 json=data,216 )217 if response.status_code != 200:218 raise RuntimeError(219 f"Sambanova /complete call failed with status code "220 f"{response.status_code}.\n Details: {response.text}"221 )222 try:223 if params.get("select_expert"):224 embedding = response.json()["predictions"]225 else:226 embedding = response.json()["predictions"]227 embeddings.extend(embedding)228 except KeyError:229 raise KeyError(230 "'predictions' not found in endpoint response",231 response.json(),232 )233 234 else:235 raise ValueError(236 f"handling of endpoint uri: {self.sambastudio_embeddings_base_uri} not implemented" # noqa: E501237 )238 239 return embeddings240 241 def embed_query(self, text: str) -> List[float]:242 """Returns a list of embeddings for the given sentences.243 Args:244 sentences (`List[str]`): List of sentences to encode245 246 Returns:247 `List[np.ndarray]` or `List[tensor]`: List of embeddings248 for the given sentences249 """250 http_session = requests.Session()251 url = self._get_full_url(252 f"{self.sambastudio_embeddings_project_id}/{self.sambastudio_embeddings_endpoint_id}"253 )254 params = json.loads(self._get_tuning_params())255 256 if "api/predict/nlp" in self.sambastudio_embeddings_base_uri:257 data = {"inputs": [text], "params": params}258 response = http_session.post(259 url,260 headers={"key": self.sambastudio_embeddings_api_key},261 json=data,262 )263 if response.status_code != 200:264 raise RuntimeError(265 f"Sambanova /complete call failed with status code "266 f"{response.status_code}.\n Details: {response.text}"267 )268 try:269 embedding = response.json()["data"][0]270 except KeyError:271 raise KeyError(272 "'data' not found in endpoint response",273 response.json(),274 )275 276 elif "api/v2/predict/generic" in self.sambastudio_embeddings_base_uri:277 data = {"items": [{"id": "item0", "value": text}], "params": params}278 response = http_session.post(279 url,280 headers={"key": self.sambastudio_embeddings_api_key},281 json=data,282 )283 if response.status_code != 200:284 raise RuntimeError(285 f"Sambanova /complete call failed with status code "286 f"{response.status_code}.\n Details: {response.text}"287 )288 try:289 embedding = response.json()["items"][0]["value"]290 except KeyError:291 raise KeyError(292 "'items' not found in endpoint response",293 response.json(),294 )295 296 elif "api/predict/generic" in self.sambastudio_embeddings_base_uri:297 data = {"instances": [text], "params": params}298 response = http_session.post(299 url,300 headers={"key": self.sambastudio_embeddings_api_key},301 json=data,302 )303 if response.status_code != 200:304 raise RuntimeError(305 f"Sambanova /complete call failed with status code "306 f"{response.status_code}.\n Details: {response.text}"307 )308 try:309 if params.get("select_expert"):310 embedding = response.json()["predictions"][0]311 else:312 embedding = response.json()["predictions"][0]313 except KeyError:314 raise KeyError(315 "'predictions' not found in endpoint response",316 response.json(),317 )318 319 else:320 raise ValueError(321 f"handling of endpoint uri: {self.sambastudio_embeddings_base_uri} not implemented" # noqa: E501322 )323 324 return embedding325 