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