codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import json4from io import StringIO5from typing import Any, Dict, Iterator, List, Optional6 7import requests8from langchain_core.callbacks.manager import CallbackManagerForLLMRun9from langchain_core.language_models.llms import LLM10from langchain_core.outputs import GenerationChunk11from langchain_core.utils import get_pydantic_field_names12from pydantic import ConfigDict13 14 15class Llamafile(LLM):16 """Llamafile lets you distribute and run large language models with a17 single file.18 19 To get started, see: https://github.com/Mozilla-Ocho/llamafile20 21 To use this class, you will need to first:22 23 1. Download a llamafile.24 2. Make the downloaded file executable: `chmod +x path/to/model.llamafile`25 3. Start the llamafile in server mode:26 27 `./path/to/model.llamafile --server --nobrowser`28 29 Example:30 .. code-block:: python31 32 from langchain_community.llms import Llamafile33 llm = Llamafile()34 llm.invoke("Tell me a joke.")35 """36 37 base_url: str = "http://localhost:8080"38 """Base url where the llamafile server is listening."""39 40 request_timeout: Optional[int] = None41 """Timeout for server requests"""42 43 streaming: bool = False44 """Allows receiving each predicted token in real-time instead of45 waiting for the completion to finish. To enable this, set to true."""46 47 # Generation options48 49 seed: int = -150 """Random Number Generator (RNG) seed. A random seed is used if this is 51 less than zero. Default: -1"""52 53 temperature: float = 0.854 """Temperature. Default: 0.8"""55 56 top_k: int = 4057 """Limit the next token selection to the K most probable tokens. 58 Default: 40."""59 60 top_p: float = 0.9561 """Limit the next token selection to a subset of tokens with a cumulative 62 probability above a threshold P. Default: 0.95."""63 64 min_p: float = 0.0565 """The minimum probability for a token to be considered, relative to 66 the probability of the most likely token. Default: 0.05."""67 68 n_predict: int = -169 """Set the maximum number of tokens to predict when generating text. 70 Note: May exceed the set limit slightly if the last token is a partial 71 multibyte character. When 0, no tokens will be generated but the prompt 72 is evaluated into the cache. Default: -1 = infinity."""73 74 n_keep: int = 075 """Specify the number of tokens from the prompt to retain when the 76 context size is exceeded and tokens need to be discarded. By default, 77 this value is set to 0 (meaning no tokens are kept). Use -1 to retain all 78 tokens from the prompt."""79 80 tfs_z: float = 1.081 """Enable tail free sampling with parameter z. Default: 1.0 = disabled."""82 83 typical_p: float = 1.084 """Enable locally typical sampling with parameter p. 85 Default: 1.0 = disabled."""86 87 repeat_penalty: float = 1.188 """Control the repetition of token sequences in the generated text. 89 Default: 1.1"""90 91 repeat_last_n: int = 6492 """Last n tokens to consider for penalizing repetition. Default: 64, 93 0 = disabled, -1 = ctx-size."""94 95 penalize_nl: bool = True96 """Penalize newline tokens when applying the repeat penalty. 97 Default: true."""98 99 presence_penalty: float = 0.0100 """Repeat alpha presence penalty. Default: 0.0 = disabled."""101 102 frequency_penalty: float = 0.0103 """Repeat alpha frequency penalty. Default: 0.0 = disabled"""104 105 mirostat: int = 0106 """Enable Mirostat sampling, controlling perplexity during text 107 generation. 0 = disabled, 1 = Mirostat, 2 = Mirostat 2.0. 108 Default: disabled."""109 110 mirostat_tau: float = 5.0111 """Set the Mirostat target entropy, parameter tau. Default: 5.0."""112 113 mirostat_eta: float = 0.1114 """Set the Mirostat learning rate, parameter eta. Default: 0.1."""115 116 model_config = ConfigDict(117 extra="forbid",118 )119 120 @property121 def _llm_type(self) -> str:122 return "llamafile"123 124 @property125 def _param_fieldnames(self) -> List[str]:126 # Return the list of fieldnames that will be passed as configurable127 # generation options to the llamafile server. Exclude 'builtin' fields128 # from the BaseLLM class like 'metadata' as well as fields that should129 # not be passed in requests (base_url, request_timeout).130 ignore_keys = [131 "base_url",132 "cache",133 "callback_manager",134 "callbacks",135 "metadata",136 "name",137 "request_timeout",138 "streaming",139 "tags",140 "verbose",141 "custom_get_token_ids",142 ]143 attrs = [144 k for k in get_pydantic_field_names(self.__class__) if k not in ignore_keys145 ]146 return attrs147 148 @property149 def _default_params(self) -> Dict[str, Any]:150 params = {}151 for fieldname in self._param_fieldnames:152 params[fieldname] = getattr(self, fieldname)153 return params154 155 def _get_parameters(156 self, stop: Optional[List[str]] = None, **kwargs: Any157 ) -> Dict[str, Any]:158 params = self._default_params159 160 # Only update keys that are already present in the default config.161 # This way, we don't accidentally post unknown/unhandled key/values162 # in the request to the llamafile server163 for k, v in kwargs.items():164 if k in params:165 params[k] = v166 167 if stop is not None and len(stop) > 0:168 params["stop"] = stop169 170 if self.streaming:171 params["stream"] = True172 173 return params174 175 def _call(176 self,177 prompt: str,178 stop: Optional[List[str]] = None,179 run_manager: Optional[CallbackManagerForLLMRun] = None,180 **kwargs: Any,181 ) -> str:182 """Request prompt completion from the llamafile server and return the183 output.184 185 Args:186 prompt: The prompt to use for generation.187 stop: A list of strings to stop generation when encountered.188 run_manager:189 **kwargs: Any additional options to pass as part of the190 generation request.191 192 Returns:193 The string generated by the model.194 195 """196 197 if self.streaming:198 with StringIO() as buff:199 for chunk in self._stream(200 prompt, stop=stop, run_manager=run_manager, **kwargs201 ):202 buff.write(chunk.text)203 204 text = buff.getvalue()205 206 return text207 208 else:209 params = self._get_parameters(stop=stop, **kwargs)210 payload = {"prompt": prompt, **params}211 212 try:213 response = requests.post(214 url=f"{self.base_url}/completion",215 headers={216 "Content-Type": "application/json",217 },218 json=payload,219 stream=False,220 timeout=self.request_timeout,221 )222 except requests.exceptions.ConnectionError:223 raise requests.exceptions.ConnectionError(224 f"Could not connect to Llamafile server. Please make sure "225 f"that a server is running at {self.base_url}."226 )227 228 response.raise_for_status()229 response.encoding = "utf-8"230 231 text = response.json()["content"]232 233 return text234 235 def _stream(236 self,237 prompt: str,238 stop: Optional[List[str]] = None,239 run_manager: Optional[CallbackManagerForLLMRun] = None,240 **kwargs: Any,241 ) -> Iterator[GenerationChunk]:242 """Yields results objects as they are generated in real time.243 244 It also calls the callback manager's on_llm_new_token event with245 similar parameters to the OpenAI LLM class method of the same name.246 247 Args:248 prompt: The prompts to pass into the model.249 stop: Optional list of stop words to use when generating.250 run_manager:251 **kwargs: Any additional options to pass as part of the252 generation request.253 254 Returns:255 A generator representing the stream of tokens being generated.256 257 Yields:258 Dictionary-like objects each containing a token259 260 Example:261 .. code-block:: python262 263 from langchain_community.llms import Llamafile264 llm = Llamafile(265 temperature = 0.0266 )267 for chunk in llm.stream("Ask 'Hi, how are you?' like a pirate:'",268 stop=["'","\n"]):269 result = chunk["choices"][0]270 print(result["text"], end='', flush=True)271 272 """273 params = self._get_parameters(stop=stop, **kwargs)274 if "stream" not in params:275 params["stream"] = True276 277 payload = {"prompt": prompt, **params}278 279 try:280 response = requests.post(281 url=f"{self.base_url}/completion",282 headers={283 "Content-Type": "application/json",284 },285 json=payload,286 stream=True,287 timeout=self.request_timeout,288 )289 except requests.exceptions.ConnectionError:290 raise requests.exceptions.ConnectionError(291 f"Could not connect to Llamafile server. Please make sure "292 f"that a server is running at {self.base_url}."293 )294 295 response.encoding = "utf8"296 297 for raw_chunk in response.iter_lines(decode_unicode=True):298 content = self._get_chunk_content(raw_chunk)299 chunk = GenerationChunk(text=content)300 301 if run_manager:302 run_manager.on_llm_new_token(token=chunk.text)303 yield chunk304 305 def _get_chunk_content(self, chunk: str) -> str:306 """When streaming is turned on, llamafile server returns lines like:307 308 'data: {"content":" They","multimodal":true,"slot_id":0,"stop":false}'309 310 Here, we convert this to a dict and return the value of the 'content'311 field312 """313 314 if chunk.startswith("data:"):315 cleaned = chunk.lstrip("data: ")316 data = json.loads(cleaned)317 return data["content"]318 else:319 return chunk320 