Team Ai
Datasetpublic

codekingpro/portable-devtools

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