Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
edenai.py268 linesDownload Raw Back to llms
1"""Wrapper around EdenAI's Generation API."""2 3import logging4from typing import Any, Dict, List, Literal, Optional5 6from aiohttp import ClientSession7from langchain_core.callbacks import (8    AsyncCallbackManagerForLLMRun,9    CallbackManagerForLLMRun,10)11from langchain_core.language_models.llms import LLM12from langchain_core.utils import get_from_dict_or_env, pre_init13from langchain_core.utils.pydantic import get_fields14from pydantic import ConfigDict, Field, model_validator15 16from langchain_community.llms.utils import enforce_stop_tokens17from langchain_community.utilities.requests import Requests18 19logger = logging.getLogger(__name__)20 21 22class EdenAI(LLM):23    """EdenAI models.24 25    To use, you should have26    the environment variable ``EDENAI_API_KEY`` set with your API token.27    You can find your token here: https://app.edenai.run/admin/account/settings28 29    `feature` and `subfeature` are required, but any other model parameters can also be30    passed in with the format params={model_param: value, ...}31 32    for api reference check edenai documentation: http://docs.edenai.co.33    """34 35    base_url: str = "https://api.edenai.run/v2"36 37    edenai_api_key: Optional[str] = None38 39    feature: Literal["text", "image"] = "text"40    """Which generative feature to use, use text by default"""41 42    subfeature: Literal["generation"] = "generation"43    """Subfeature of above feature, use generation by default"""44 45    provider: str46    """Generative provider to use (eg: openai,stabilityai,cohere,google etc.)"""47 48    model: Optional[str] = None49    """50    model name for above provider (eg: 'gpt-3.5-turbo-instruct' for openai)51    available models are shown on https://docs.edenai.co/ under 'available providers'52    """53 54    # Optional parameters to add depending of chosen feature55    # see api reference for more infos56    temperature: Optional[float] = Field(default=None, ge=0, le=1)  # for text57    max_tokens: Optional[int] = Field(default=None, ge=0)  # for text58    resolution: Optional[Literal["256x256", "512x512", "1024x1024"]] = None  # for image59 60    params: Dict[str, Any] = Field(default_factory=dict)61    """62    DEPRECATED: use temperature, max_tokens, resolution directly63    optional parameters to pass to api64    """65 66    model_kwargs: Dict[str, Any] = Field(default_factory=dict)67    """extra parameters"""68 69    stop_sequences: Optional[List[str]] = None70    """Stop sequences to use."""71 72    model_config = ConfigDict(73        extra="forbid",74    )75 76    @pre_init77    def validate_environment(cls, values: Dict) -> Dict:78        """Validate that api key exists in environment."""79        values["edenai_api_key"] = get_from_dict_or_env(80            values, "edenai_api_key", "EDENAI_API_KEY"81        )82        return values83 84    @model_validator(mode="before")85    @classmethod86    def build_extra(cls, values: Dict[str, Any]) -> Any:87        """Build extra kwargs from additional params that were passed in."""88        all_required_field_names = {field.alias for field in get_fields(cls).values()}89 90        extra = values.get("model_kwargs", {})91        for field_name in list(values):92            if field_name not in all_required_field_names:93                if field_name in extra:94                    raise ValueError(f"Found {field_name} supplied twice.")95                logger.warning(96                    f"""{field_name} was transferred to model_kwargs.97                    Please confirm that {field_name} is what you intended."""98                )99                extra[field_name] = values.pop(field_name)100        values["model_kwargs"] = extra101        return values102 103    @property104    def _llm_type(self) -> str:105        """Return type of model."""106        return "edenai"107 108    def _format_output(self, output: dict) -> str:109        if self.feature == "text":110            return output[self.provider]["generated_text"]111        else:112            return output[self.provider]["items"][0]["image"]113 114    @staticmethod115    def get_user_agent() -> str:116        from langchain_community import __version__117 118        return f"langchain/{__version__}"119 120    def _call(121        self,122        prompt: str,123        stop: Optional[List[str]] = None,124        run_manager: Optional[CallbackManagerForLLMRun] = None,125        **kwargs: Any,126    ) -> str:127        """Call out to EdenAI's text generation endpoint.128 129        Args:130            prompt: The prompt to pass into the model.131 132        Returns:133            json formatted str response.134        """135        stops = None136        if self.stop_sequences is not None and stop is not None:137            raise ValueError(138                "stop sequences found in both the input and default params."139            )140        elif self.stop_sequences is not None:141            stops = self.stop_sequences142        else:143            stops = stop144 145        url = f"{self.base_url}/{self.feature}/{self.subfeature}"146        headers = {147            "Authorization": f"Bearer {self.edenai_api_key}",148            "User-Agent": self.get_user_agent(),149        }150        payload: Dict[str, Any] = {151            "providers": self.provider,152            "text": prompt,153            "max_tokens": self.max_tokens,154            "temperature": self.temperature,155            "resolution": self.resolution,156            **self.params,157            **kwargs,158            "num_images": 1,  # always limit to 1 (ignored for text)159        }160 161        # filter None values to not pass them to the http payload162        payload = {k: v for k, v in payload.items() if v is not None}163 164        if self.model is not None:165            payload["settings"] = {self.provider: self.model}166 167        request = Requests(headers=headers)168        response = request.post(url=url, data=payload)169 170        if response.status_code >= 500:171            raise Exception(f"EdenAI Server: Error {response.status_code}")172        elif response.status_code >= 400:173            raise ValueError(f"EdenAI received an invalid payload: {response.text}")174        elif response.status_code != 200:175            raise Exception(176                f"EdenAI returned an unexpected response with status "177                f"{response.status_code}: {response.text}"178            )179 180        data = response.json()181        provider_response = data[self.provider]182        if provider_response.get("status") == "fail":183            err_msg = provider_response.get("error", {}).get("message")184            raise Exception(err_msg)185 186        output = self._format_output(data)187 188        if stops is not None:189            output = enforce_stop_tokens(output, stops)190 191        return output192 193    async def _acall(194        self,195        prompt: str,196        stop: Optional[List[str]] = None,197        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,198        **kwargs: Any,199    ) -> str:200        """Call EdenAi model to get predictions based on the prompt.201 202        Args:203            prompt: The prompt to pass into the model.204            stop: A list of stop words (optional).205            run_manager: A callback manager for async interaction with LLMs.206 207        Returns:208            The string generated by the model.209        """210 211        stops = None212        if self.stop_sequences is not None and stop is not None:213            raise ValueError(214                "stop sequences found in both the input and default params."215            )216        elif self.stop_sequences is not None:217            stops = self.stop_sequences218        else:219            stops = stop220 221        url = f"{self.base_url}/{self.feature}/{self.subfeature}"222        headers = {223            "Authorization": f"Bearer {self.edenai_api_key}",224            "User-Agent": self.get_user_agent(),225        }226        payload: Dict[str, Any] = {227            "providers": self.provider,228            "text": prompt,229            "max_tokens": self.max_tokens,230            "temperature": self.temperature,231            "resolution": self.resolution,232            **self.params,233            **kwargs,234            "num_images": 1,  # always limit to 1 (ignored for text)235        }236 237        # filter `None` values to not pass them to the http payload as null238        payload = {k: v for k, v in payload.items() if v is not None}239 240        if self.model is not None:241            payload["settings"] = {self.provider: self.model}242 243        async with ClientSession() as session:244            async with session.post(url, json=payload, headers=headers) as response:245                if response.status >= 500:246                    raise Exception(f"EdenAI Server: Error {response.status}")247                elif response.status >= 400:248                    raise ValueError(249                        f"EdenAI received an invalid payload: {response.text}"250                    )251                elif response.status != 200:252                    raise Exception(253                        f"EdenAI returned an unexpected response with status "254                        f"{response.status}: {response.text}"255                    )256 257                response_json = await response.json()258                provider_response = response_json[self.provider]259                if provider_response.get("status") == "fail":260                    err_msg = provider_response.get("error", {}).get("message")261                    raise Exception(err_msg)262 263                output = self._format_output(response_json)264                if stops is not None:265                    output = enforce_stop_tokens(output, stops)266 267                return output268 
codekingpro/portable-devtools · Team Ai