Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
writer.py198 linesDownload Raw Back to llms
1from typing import Any, AsyncIterator, Dict, Iterator, List, Mapping, Optional2 3from langchain_core.callbacks import (4    AsyncCallbackManagerForLLMRun,5    CallbackManagerForLLMRun,6)7from langchain_core.language_models.llms import LLM8from langchain_core.outputs import GenerationChunk9from langchain_core.utils import get_from_dict_or_env10from pydantic import ConfigDict, Field, SecretStr, model_validator11 12 13class Writer(LLM):14    """Writer large language models.15 16    To use, you should have the ``writer-sdk`` Python package installed, and the17    environment variable ``WRITER_API_KEY`` set with your API key.18 19    Example:20        .. code-block:: python21 22            from langchain_community.llms import Writer as WriterLLM23            from writerai import Writer, AsyncWriter24 25            client = Writer()26            async_client = AsyncWriter()27 28            chat = WriterLLM(29                client=client,30                async_client=async_client31            )32    """33 34    client: Any = Field(default=None, exclude=True)  #: :meta private:35    async_client: Any = Field(default=None, exclude=True)  #: :meta private:36 37    api_key: Optional[SecretStr] = Field(default=None)38    """Writer API key."""39 40    model_name: str = Field(default="palmyra-x-003-instruct", alias="model")41    """Model name to use."""42 43    max_tokens: Optional[int] = None44    """The maximum number of tokens that the model can generate in the response."""45 46    temperature: Optional[float] = 0.747    """Controls the randomness of the model's outputs. Higher values lead to more 48    random outputs, while lower values make the model more deterministic."""49 50    top_p: Optional[float] = None51    """Used to control the nucleus sampling, where only the most probable tokens52     with a cumulative probability of top_p are considered for sampling, providing 53     a way to fine-tune the randomness of predictions."""54 55    stop: Optional[List[str]] = None56    """Specifies stopping conditions for the model's output generation. This can57     be an array of strings or a single string that the model will look for as a 58     signal to stop generating further tokens."""59 60    best_of: Optional[int] = None61    """Specifies the number of completions to generate and return the best one.62     Useful for generating multiple outputs and choosing the best based on some63      criteria."""64 65    model_kwargs: Dict[str, Any] = Field(default_factory=dict)66    """Holds any model parameters valid for `create` call not explicitly specified."""67 68    model_config = ConfigDict(populate_by_name=True)69 70    @property71    def _default_params(self) -> Mapping[str, Any]:72        """Get the default parameters for calling Writer API."""73        return {74            "max_tokens": self.max_tokens,75            "temperature": self.temperature,76            "top_p": self.top_p,77            "stop": self.stop,78            "best_of": self.best_of,79            **self.model_kwargs,80        }81 82    @property83    def _identifying_params(self) -> Mapping[str, Any]:84        """Get the identifying parameters."""85        return {86            "model": self.model_name,87            **self._default_params,88        }89 90    @property91    def _llm_type(self) -> str:92        """Return type of llm."""93        return "writer"94 95    @model_validator(mode="before")96    @classmethod97    def validate_environment(cls, values: Dict) -> Any:98        """Validates that api key is passed and creates Writer clients."""99        try:100            from writerai import AsyncClient, Client101        except ImportError as e:102            raise ImportError(103                "Could not import writerai python package. "104                "Please install it with `pip install writerai`."105            ) from e106 107        if not values.get("client"):108            values.update(109                {110                    "client": Client(111                        api_key=get_from_dict_or_env(112                            values, "api_key", "WRITER_API_KEY"113                        )114                    )115                }116            )117 118        if not values.get("async_client"):119            values.update(120                {121                    "async_client": AsyncClient(122                        api_key=get_from_dict_or_env(123                            values, "api_key", "WRITER_API_KEY"124                        )125                    )126                }127            )128 129        if not (130            type(values.get("client")) is Client131            and type(values.get("async_client")) is AsyncClient132        ):133            raise ValueError(134                "'client' attribute must be with type 'Client' and "135                "'async_client' must be with type 'AsyncClient' from 'writerai' package"136            )137 138        return values139 140    def _call(141        self,142        prompt: str,143        stop: Optional[List[str]] = None,144        run_manager: Optional[CallbackManagerForLLMRun] = None,145        **kwargs: Any,146    ) -> str:147        params = {**self._identifying_params, **kwargs}148        if stop is not None:149            params.update({"stop": stop})150        text = self.client.completions.create(prompt=prompt, **params).choices[0].text151        return text152 153    async def _acall(154        self,155        prompt: str,156        stop: Optional[list[str]] = None,157        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,158        **kwargs: Any,159    ) -> str:160        params = {**self._identifying_params, **kwargs}161        if stop is not None:162            params.update({"stop": stop})163        response = await self.async_client.completions.create(prompt=prompt, **params)164        text = response.choices[0].text165        return text166 167    def _stream(168        self,169        prompt: str,170        stop: Optional[list[str]] = None,171        run_manager: Optional[CallbackManagerForLLMRun] = None,172        **kwargs: Any,173    ) -> Iterator[GenerationChunk]:174        params = {**self._identifying_params, **kwargs, "stream": True}175        if stop is not None:176            params.update({"stop": stop})177        response = self.client.completions.create(prompt=prompt, **params)178        for chunk in response:179            if run_manager:180                run_manager.on_llm_new_token(chunk.value)181            yield GenerationChunk(text=chunk.value)182 183    async def _astream(184        self,185        prompt: str,186        stop: Optional[list[str]] = None,187        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,188        **kwargs: Any,189    ) -> AsyncIterator[GenerationChunk]:190        params = {**self._identifying_params, **kwargs, "stream": True}191        if stop is not None:192            params.update({"stop": stop})193        response = await self.async_client.completions.create(prompt=prompt, **params)194        async for chunk in response:195            if run_manager:196                await run_manager.on_llm_new_token(chunk.value)197            yield GenerationChunk(text=chunk.value)198 
codekingpro/portable-devtools · Team Ai