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