Team Ai
Datasetpublic

codekingpro/portable-devtools

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