Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
deepinfra.py247 linesDownload Raw Back to llms
1import json2from typing import Any, AsyncIterator, Dict, Iterator, List, Mapping, Optional3 4import aiohttp5from langchain_core.callbacks import (6    AsyncCallbackManagerForLLMRun,7    CallbackManagerForLLMRun,8)9from langchain_core.language_models.llms import LLM10from langchain_core.outputs import GenerationChunk11from langchain_core.utils import get_from_dict_or_env, pre_init12from pydantic import ConfigDict13 14from langchain_community.utilities.requests import Requests15 16DEFAULT_MODEL_ID = "meta-llama/Meta-Llama-3-70B-Instruct"17 18 19class DeepInfra(LLM):20    """DeepInfra models.21 22    To use, you should have the environment variable ``DEEPINFRA_API_TOKEN``23    set with your API token, or pass it as a named parameter to the24    constructor.25 26    Only supports `text-generation` and `text2text-generation` for now.27 28    Example:29        .. code-block:: python30 31            from langchain_community.llms import DeepInfra32            di = DeepInfra(model_id="google/flan-t5-xl",33                                deepinfra_api_token="my-api-key")34    """35 36    model_id: str = DEFAULT_MODEL_ID37    model_kwargs: Optional[Dict] = None38 39    deepinfra_api_token: Optional[str] = None40 41    model_config = ConfigDict(42        extra="forbid",43    )44 45    @pre_init46    def validate_environment(cls, values: Dict) -> Dict:47        """Validate that api key and python package exists in environment."""48        deepinfra_api_token = get_from_dict_or_env(49            values, "deepinfra_api_token", "DEEPINFRA_API_TOKEN"50        )51        values["deepinfra_api_token"] = deepinfra_api_token52        return values53 54    @property55    def _identifying_params(self) -> Mapping[str, Any]:56        """Get the identifying parameters."""57        return {58            **{"model_id": self.model_id},59            **{"model_kwargs": self.model_kwargs},60        }61 62    @property63    def _llm_type(self) -> str:64        """Return type of llm."""65        return "deepinfra"66 67    def _url(self) -> str:68        return f"https://api.deepinfra.com/v1/inference/{self.model_id}"69 70    def _headers(self) -> Dict:71        return {72            "Authorization": f"bearer {self.deepinfra_api_token}",73            "Content-Type": "application/json",74        }75 76    def _body(self, prompt: str, kwargs: Any) -> Dict:77        model_kwargs = self.model_kwargs or {}78        model_kwargs = {**model_kwargs, **kwargs}79 80        return {81            "input": prompt,82            **model_kwargs,83        }84 85    def _handle_status(self, code: int, text: Any) -> None:86        if code >= 500:87            raise Exception(f"DeepInfra Server: Error {text}")88        elif code == 401:89            raise Exception("DeepInfra Server: Unauthorized")90        elif code == 403:91            raise Exception("DeepInfra Server: Unauthorized")92        elif code == 404:93            raise Exception(f"DeepInfra Server: Model not found {self.model_id}")94        elif code == 429:95            raise Exception("DeepInfra Server: Rate limit exceeded")96        elif code >= 400:97            raise ValueError(f"DeepInfra received an invalid payload: {text}")98        elif code != 200:99            raise Exception(100                f"DeepInfra returned an unexpected response with status {code}: {text}"101            )102 103    def _call(104        self,105        prompt: str,106        stop: Optional[List[str]] = None,107        run_manager: Optional[CallbackManagerForLLMRun] = None,108        **kwargs: Any,109    ) -> str:110        """Call out to DeepInfra's inference API endpoint.111 112        Args:113            prompt: The prompt to pass into the model.114            stop: Optional list of stop words to use when generating.115 116        Returns:117            The string generated by the model.118 119        Example:120            .. code-block:: python121 122                response = di("Tell me a joke.")123        """124 125        request = Requests(headers=self._headers())126        response = request.post(url=self._url(), data=self._body(prompt, kwargs))127 128        self._handle_status(response.status_code, response.text)129        data = response.json()130 131        return data["results"][0]["generated_text"]132 133    async def _acall(134        self,135        prompt: str,136        stop: Optional[List[str]] = None,137        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,138        **kwargs: Any,139    ) -> str:140        request = Requests(headers=self._headers())141        async with request.apost(142            url=self._url(), data=self._body(prompt, kwargs)143        ) as response:144            self._handle_status(response.status, response.text)145            data = await response.json()146            return data["results"][0]["generated_text"]147 148    def _stream(149        self,150        prompt: str,151        stop: Optional[List[str]] = None,152        run_manager: Optional[CallbackManagerForLLMRun] = None,153        **kwargs: Any,154    ) -> Iterator[GenerationChunk]:155        request = Requests(headers=self._headers())156        response = request.post(157            url=self._url(), data=self._body(prompt, {**kwargs, "stream": True})158        )159        response_text = response.text160        self._handle_body_errors(response_text)161        self._handle_status(response.status_code, response.text)162        for line in _parse_stream(response.iter_lines()):163            chunk = _handle_sse_line(line)164            if chunk:165                if run_manager:166                    run_manager.on_llm_new_token(chunk.text)167                yield chunk168 169    async def _astream(170        self,171        prompt: str,172        stop: Optional[List[str]] = None,173        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,174        **kwargs: Any,175    ) -> AsyncIterator[GenerationChunk]:176        request = Requests(headers=self._headers())177        async with request.apost(178            url=self._url(), data=self._body(prompt, {**kwargs, "stream": True})179        ) as response:180            response_text = await response.text()181            self._handle_body_errors(response_text)182            self._handle_status(response.status, response.text)183            async for line in _parse_stream_async(response.content):184                chunk = _handle_sse_line(line)185                if chunk:186                    if run_manager:187                        await run_manager.on_llm_new_token(chunk.text)188                    yield chunk189 190    def _handle_body_errors(self, body: str) -> None:191        """192        Example error response:193        data: {"error_type": "validation_error",194        "error_message": "ConnectionError: ..."}195        """196        if "error" in body:197            try:198                # Remove data: prefix if present199                if body.startswith("data:"):200                    body = body[len("data:") :]201                error_data = json.loads(body)202                error_message = error_data.get("error_message", "Unknown error")203 204                raise Exception(f"DeepInfra Server Error: {error_message}")205            except json.JSONDecodeError:206                raise Exception(f"DeepInfra Server: {body}")207 208 209def _parse_stream(rbody: Iterator[bytes]) -> Iterator[str]:210    for line in rbody:211        _line = _parse_stream_helper(line)212        if _line is not None:213            yield _line214 215 216async def _parse_stream_async(rbody: aiohttp.StreamReader) -> AsyncIterator[str]:217    async for line in rbody:218        _line = _parse_stream_helper(line)219        if _line is not None:220            yield _line221 222 223def _parse_stream_helper(line: bytes) -> Optional[str]:224    if line and line.startswith(b"data:"):225        if line.startswith(b"data: "):226            # SSE event may be valid when it contain whitespace227            line = line[len(b"data: ") :]228        else:229            line = line[len(b"data:") :]230        if line.strip() == b"[DONE]":231            # return here will cause GeneratorExit exception in urllib3232            # and it will close http connection with TCP Reset233            return None234        else:235            return line.decode("utf-8")236    return None237 238 239def _handle_sse_line(line: str) -> Optional[GenerationChunk]:240    try:241        obj = json.loads(line)242        return GenerationChunk(243            text=obj.get("token", {}).get("text"),244        )245    except Exception:246        return None247 
codekingpro/portable-devtools · Team Ai