codekingpro/portable-devtools
114k
1import logging2from typing import Any, Dict, List, Mapping, Optional3 4from langchain_core.callbacks import CallbackManagerForLLMRun5from langchain_core.language_models.llms import LLM6from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init7from langchain_core.utils.pydantic import get_fields8from pydantic import ConfigDict, Field, SecretStr, model_validator9 10from langchain_community.llms.utils import enforce_stop_tokens11 12logger = logging.getLogger(__name__)13 14 15class Petals(LLM):16 """Petals Bloom models.17 18 To use, you should have the ``petals`` python package installed, and the19 environment variable ``HUGGINGFACE_API_KEY`` set with your API key.20 21 Any parameters that are valid to be passed to the call can be passed22 in, even if not explicitly saved on this class.23 24 Example:25 .. code-block:: python26 27 from langchain_community.llms import petals28 petals = Petals()29 30 """31 32 client: Any = None33 """The client to use for the API calls."""34 35 tokenizer: Any = None36 """The tokenizer to use for the API calls."""37 38 model_name: str = "bigscience/bloom-petals"39 """The model to use."""40 41 temperature: float = 0.742 """What sampling temperature to use"""43 44 max_new_tokens: int = 25645 """The maximum number of new tokens to generate in the completion."""46 47 top_p: float = 0.948 """The cumulative probability for top-p sampling."""49 50 top_k: Optional[int] = None51 """The number of highest probability vocabulary tokens52 to keep for top-k-filtering."""53 54 do_sample: bool = True55 """Whether or not to use sampling; use greedy decoding otherwise."""56 57 max_length: Optional[int] = None58 """The maximum length of the sequence to be generated."""59 60 model_kwargs: Dict[str, Any] = Field(default_factory=dict)61 """Holds any model parameters valid for `create` call62 not explicitly specified."""63 64 huggingface_api_key: Optional[SecretStr] = None65 66 model_config = ConfigDict(67 extra="forbid",68 )69 70 @model_validator(mode="before")71 @classmethod72 def build_extra(cls, values: Dict[str, Any]) -> Any:73 """Build extra kwargs from additional params that were passed in."""74 all_required_field_names = {field.alias for field in get_fields(cls).values()}75 76 extra = values.get("model_kwargs", {})77 for field_name in list(values):78 if field_name not in all_required_field_names:79 if field_name in extra:80 raise ValueError(f"Found {field_name} supplied twice.")81 logger.warning(82 f"""WARNING! {field_name} is not default parameter.83 {field_name} was transferred to model_kwargs.84 Please confirm that {field_name} is what you intended."""85 )86 extra[field_name] = values.pop(field_name)87 values["model_kwargs"] = extra88 return values89 90 @pre_init91 def validate_environment(cls, values: Dict) -> Dict:92 """Validate that api key and python package exists in environment."""93 huggingface_api_key = convert_to_secret_str(94 get_from_dict_or_env(values, "huggingface_api_key", "HUGGINGFACE_API_KEY")95 )96 try:97 from petals import AutoDistributedModelForCausalLM98 from transformers import AutoTokenizer99 100 model_name = values["model_name"]101 values["tokenizer"] = AutoTokenizer.from_pretrained(model_name)102 values["client"] = AutoDistributedModelForCausalLM.from_pretrained(103 model_name104 )105 values["huggingface_api_key"] = huggingface_api_key.get_secret_value()106 107 except ImportError:108 raise ImportError(109 "Could not import transformers or petals python package."110 "Please install with `pip install -U transformers petals`."111 )112 return values113 114 @property115 def _default_params(self) -> Dict[str, Any]:116 """Get the default parameters for calling Petals API."""117 normal_params = {118 "temperature": self.temperature,119 "max_new_tokens": self.max_new_tokens,120 "top_p": self.top_p,121 "top_k": self.top_k,122 "do_sample": self.do_sample,123 "max_length": self.max_length,124 }125 return {**normal_params, **self.model_kwargs}126 127 @property128 def _identifying_params(self) -> Mapping[str, Any]:129 """Get the identifying parameters."""130 return {**{"model_name": self.model_name}, **self._default_params}131 132 @property133 def _llm_type(self) -> str:134 """Return type of llm."""135 return "petals"136 137 def _call(138 self,139 prompt: str,140 stop: Optional[List[str]] = None,141 run_manager: Optional[CallbackManagerForLLMRun] = None,142 **kwargs: Any,143 ) -> str:144 """Call the Petals API."""145 params = self._default_params146 params = {**params, **kwargs}147 inputs = self.tokenizer(prompt, return_tensors="pt")["input_ids"]148 outputs = self.client.generate(inputs, **params)149 text = self.tokenizer.decode(outputs[0])150 if stop is not None:151 # I believe this is required since the stop tokens152 # are not enforced by the model parameters153 text = enforce_stop_tokens(text, stop)154 return text155 