Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
predibase.py219 linesDownload Raw Back to llms
1import os2from typing import Any, Dict, List, Mapping, Optional, Union3 4from langchain_core.callbacks import CallbackManagerForLLMRun5from langchain_core.language_models.llms import LLM6from pydantic import Field, SecretStr7 8 9class Predibase(LLM):10    """Use your Predibase models with Langchain.11 12    To use, you should have the ``predibase`` python package installed,13    and have your Predibase API key.14 15    The `model` parameter is the Predibase "serverless" base_model ID16    (see https://docs.predibase.com/user-guide/inference/models for the catalog).17 18    An optional `adapter_id` parameter is the Predibase ID or HuggingFace ID of a19    fine-tuned LLM adapter, whose base model is the `model` parameter; the20    fine-tuned adapter must be compatible with its base model;21    otherwise, an error is raised.  If the fine-tuned adapter is hosted at Predibase,22    then `adapter_version` in the adapter repository must be specified.23 24    An optional `predibase_sdk_version` parameter defaults to latest SDK version.25    """26 27    model: str28    predibase_api_key: SecretStr29    predibase_sdk_version: Optional[str] = None30    adapter_id: Optional[str] = None31    adapter_version: Optional[int] = None32    model_kwargs: Dict[str, Any] = Field(default_factory=dict)33    default_options_for_generation: dict = Field(34        {35            "max_new_tokens": 256,36            "temperature": 0.1,37        }38    )39 40    @property41    def _llm_type(self) -> str:42        return "predibase"43 44    def _call(45        self,46        prompt: str,47        stop: Optional[List[str]] = None,48        run_manager: Optional[CallbackManagerForLLMRun] = None,49        **kwargs: Any,50    ) -> str:51        options: Dict[str, Union[str, float]] = {52            **self.default_options_for_generation,53            **(self.model_kwargs or {}),54            **(kwargs or {}),55        }56        if self._is_deprecated_sdk_version():57            try:58                from predibase import PredibaseClient59                from predibase.pql import get_session60                from predibase.pql.api import (61                    ServerResponseError,62                    Session,63                )64                from predibase.resource.llm.interface import (65                    HuggingFaceLLM,66                    LLMDeployment,67                )68                from predibase.resource.llm.response import GeneratedResponse69                from predibase.resource.model import Model70 71                session: Session = get_session(72                    token=self.predibase_api_key.get_secret_value(),73                    gateway="https://api.app.predibase.com/v1",74                    serving_endpoint="serving.app.predibase.com",75                )76                pc: PredibaseClient = PredibaseClient(session=session)77            except ImportError as e:78                raise ImportError(79                    "Could not import Predibase Python package. "80                    "Please install it with `pip install predibase`."81                ) from e82            except ValueError as e:83                raise ValueError("Your API key is not correct. Please try again") from e84 85            base_llm_deployment: LLMDeployment = pc.LLM(86                uri=f"pb://deployments/{self.model}"87            )88            result: GeneratedResponse89            if self.adapter_id:90                """91                Attempt to retrieve the fine-tuned adapter from a Predibase92                repository.  If absent, then load the fine-tuned adapter93                from a HuggingFace repository.94                """95                adapter_model: Union[Model, HuggingFaceLLM]96                try:97                    adapter_model = pc.get_model(98                        name=self.adapter_id,99                        version=self.adapter_version,100                        model_id=None,101                    )102                except ServerResponseError:103                    # Predibase does not recognize the adapter ID (query HuggingFace).104                    adapter_model = pc.LLM(uri=f"hf://{self.adapter_id}")105                result = base_llm_deployment.with_adapter(model=adapter_model).generate(106                    prompt=prompt,107                    options=options,108                )109            else:110                result = base_llm_deployment.generate(111                    prompt=prompt,112                    options=options,113                )114            return result.response115 116        from predibase import Predibase117 118        os.environ["PREDIBASE_GATEWAY"] = "https://api.app.predibase.com"119        predibase: Predibase = Predibase(120            api_token=self.predibase_api_key.get_secret_value()121        )122 123        import requests124        from lorax.client import Client as LoraxClient125        from lorax.errors import GenerationError126        from lorax.types import Response127 128        lorax_client: LoraxClient = predibase.deployments.client(129            deployment_ref=self.model130        )131 132        response: Response133        if self.adapter_id:134            """135            Attempt to retrieve the fine-tuned adapter from a Predibase repository.136            If absent, then load the fine-tuned adapter from a HuggingFace repository.137            """138            if self.adapter_version:139                # Since the adapter version is provided, query the Predibase repository.140                pb_adapter_id: str = f"{self.adapter_id}/{self.adapter_version}"141                options.pop(142                    "api_token", None143                )  # The "api_token" is not used for Predibase-hosted models.144                try:145                    response = lorax_client.generate(146                        prompt=prompt,147                        adapter_id=pb_adapter_id,148                        **options,149                    )150                except GenerationError as ge:151                    raise ValueError(152                        f"""An adapter with the ID "{pb_adapter_id}" cannot be \153found in the Predibase repository of fine-tuned adapters."""154                    ) from ge155            else:156                # The adapter version is omitted,157                # hence look for the adapter ID in the HuggingFace repository.158                try:159                    response = lorax_client.generate(160                        prompt=prompt,161                        adapter_id=self.adapter_id,162                        adapter_source="hub",163                        **options,164                    )165                except GenerationError as ge:166                    raise ValueError(167                        f"""Either an adapter with the ID "{self.adapter_id}" \168cannot be found in a HuggingFace repository, or it is incompatible with the \169base model (please make sure that the adapter configuration is consistent).170"""171                    ) from ge172        else:173            try:174                response = lorax_client.generate(175                    prompt=prompt,176                    **options,177                )178            except requests.JSONDecodeError as jde:179                raise ValueError(180                    f"""An LLM with the deployment ID "{self.model}" cannot be found \181at Predibase (please refer to \182"https://docs.predibase.com/user-guide/inference/models" for the list of \183supported models).184"""185                ) from jde186        response_text = response.generated_text187 188        return response_text189 190    @property191    def _identifying_params(self) -> Mapping[str, Any]:192        """Get the identifying parameters."""193        return {194            **{"model_kwargs": self.model_kwargs},195        }196 197    def _is_deprecated_sdk_version(self) -> bool:198        try:199            import semantic_version200            from predibase.version import __version__ as current_version201            from semantic_version.base import Version202 203            sdk_semver_deprecated: Version = semantic_version.Version(204                version_string="2024.4.8"205            )206            actual_current_version: str = self.predibase_sdk_version or current_version207            sdk_semver_current: Version = semantic_version.Version(208                version_string=actual_current_version209            )210            return not (211                (sdk_semver_current > sdk_semver_deprecated)212                or ("+dev" in actual_current_version)213            )214        except ImportError as e:215            raise ImportError(216                "Could not import Predibase Python package. "217                "Please install it with `pip install semantic_version predibase`."218            ) from e219 
codekingpro/portable-devtools · Team Ai