Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
modal.py102 linesDownload Raw Back to llms
1import logging2from typing import Any, Dict, List, Mapping, Optional3 4import requests5from langchain_core.callbacks import CallbackManagerForLLMRun6from langchain_core.language_models.llms import LLM7from langchain_core.utils.pydantic import get_fields8from pydantic import ConfigDict, Field, model_validator9 10from langchain_community.llms.utils import enforce_stop_tokens11 12logger = logging.getLogger(__name__)13 14 15class Modal(LLM):16    """Modal large language models.17 18    To use, you should have the ``modal-client`` python package installed.19 20    Any parameters that are valid to be passed to the call can be passed21    in, even if not explicitly saved on this class.22 23    Example:24        .. code-block:: python25 26            from langchain_community.llms import Modal27            modal = Modal(endpoint_url="")28 29    """30 31    endpoint_url: str = ""32    """model endpoint to use"""33 34    model_kwargs: Dict[str, Any] = Field(default_factory=dict)35    """Holds any model parameters valid for `create` call not36    explicitly specified."""37 38    model_config = ConfigDict(39        extra="forbid",40    )41 42    @model_validator(mode="before")43    @classmethod44    def build_extra(cls, values: Dict[str, Any]) -> Any:45        """Build extra kwargs from additional params that were passed in."""46        all_required_field_names = {field.alias for field in get_fields(cls).values()}47 48        extra = values.get("model_kwargs", {})49        for field_name in list(values):50            if field_name not in all_required_field_names:51                if field_name in extra:52                    raise ValueError(f"Found {field_name} supplied twice.")53                logger.warning(54                    f"""{field_name} was transferred to model_kwargs.55                    Please confirm that {field_name} is what you intended."""56                )57                extra[field_name] = values.pop(field_name)58        values["model_kwargs"] = extra59        return values60 61    @property62    def _identifying_params(self) -> Mapping[str, Any]:63        """Get the identifying parameters."""64        return {65            **{"endpoint_url": self.endpoint_url},66            **{"model_kwargs": self.model_kwargs},67        }68 69    @property70    def _llm_type(self) -> str:71        """Return type of llm."""72        return "modal"73 74    def _call(75        self,76        prompt: str,77        stop: Optional[List[str]] = None,78        run_manager: Optional[CallbackManagerForLLMRun] = None,79        **kwargs: Any,80    ) -> str:81        """Call to Modal endpoint."""82        params = self.model_kwargs or {}83        params = {**params, **kwargs}84        response = requests.post(85            url=self.endpoint_url,86            headers={87                "Content-Type": "application/json",88            },89            json={"prompt": prompt, **params},90        )91        try:92            if prompt in response.json()["prompt"]:93                response_json = response.json()94        except KeyError:95            raise KeyError("LangChain requires 'prompt' key in response.")96        text = response_json["prompt"]97        if stop is not None:98            # I believe this is required since the stop tokens99            # are not enforced by the model parameters100            text = enforce_stop_tokens(text, stop)101        return text102 
codekingpro/portable-devtools · Team Ai