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