codekingpro/portable-devtools
114k
1import logging2from typing import Any, Dict, List, Optional3 4import requests5from langchain_core.callbacks import CallbackManagerForLLMRun6from langchain_core.language_models.llms import LLM7 8logger = logging.getLogger(__name__)9 10 11def clean_url(url: str) -> str:12 """Remove trailing slash and /api from url if present."""13 if url.endswith("/api"):14 return url[:-4]15 elif url.endswith("/"):16 return url[:-1]17 else:18 return url19 20 21class KoboldApiLLM(LLM):22 """Kobold API language model.23 24 It includes several fields that can be used to control the text generation process.25 26 To use this class, instantiate it with the required parameters and call it with a27 prompt to generate text. For example:28 29 kobold = KoboldApiLLM(endpoint="http://localhost:5000")30 result = kobold("Write a story about a dragon.")31 32 This will send a POST request to the Kobold API with the provided prompt and33 generate text.34 """35 36 endpoint: str37 """The API endpoint to use for generating text."""38 39 use_story: Optional[bool] = False40 """ Whether or not to use the story from the KoboldAI GUI when generating text. """41 42 use_authors_note: Optional[bool] = False43 """Whether to use the author's note from the KoboldAI GUI when generating text.44 45 This has no effect unless use_story is also enabled.46 """47 48 use_world_info: Optional[bool] = False49 """Whether to use the world info from the KoboldAI GUI when generating text."""50 51 use_memory: Optional[bool] = False52 """Whether to use the memory from the KoboldAI GUI when generating text."""53 54 max_context_length: Optional[int] = 160055 """Maximum number of tokens to send to the model.56 57 minimum: 158 """59 60 max_length: Optional[int] = 8061 """Number of tokens to generate.62 63 maximum: 51264 minimum: 165 """66 67 rep_pen: Optional[float] = 1.1268 """Base repetition penalty value.69 70 minimum: 171 """72 73 rep_pen_range: Optional[int] = 102474 """Repetition penalty range.75 76 minimum: 077 """78 79 rep_pen_slope: Optional[float] = 0.980 """Repetition penalty slope.81 82 minimum: 083 """84 85 temperature: Optional[float] = 0.686 """Temperature value.87 88 exclusiveMinimum: 089 """90 91 tfs: Optional[float] = 0.992 """Tail free sampling value.93 94 maximum: 195 minimum: 096 """97 98 top_a: Optional[float] = 0.999 """Top-a sampling value.100 101 minimum: 0102 """103 104 top_p: Optional[float] = 0.95105 """Top-p sampling value.106 107 maximum: 1108 minimum: 0109 """110 111 top_k: Optional[int] = 0112 """Top-k sampling value.113 114 minimum: 0115 """116 117 typical: Optional[float] = 0.5118 """Typical sampling value.119 120 maximum: 1121 minimum: 0122 """123 124 @property125 def _llm_type(self) -> str:126 return "koboldai"127 128 def _call(129 self,130 prompt: str,131 stop: Optional[List[str]] = None,132 run_manager: Optional[CallbackManagerForLLMRun] = None,133 **kwargs: Any,134 ) -> str:135 """Call the API and return the output.136 137 Args:138 prompt: The prompt to use for generation.139 stop: A list of strings to stop generation when encountered.140 141 Returns:142 The generated text.143 144 Example:145 .. code-block:: python146 147 from langchain_community.llms import KoboldApiLLM148 149 llm = KoboldApiLLM(endpoint="http://localhost:5000")150 llm.invoke("Write a story about dragons.")151 """152 data: Dict[str, Any] = {153 "prompt": prompt,154 "use_story": self.use_story,155 "use_authors_note": self.use_authors_note,156 "use_world_info": self.use_world_info,157 "use_memory": self.use_memory,158 "max_context_length": self.max_context_length,159 "max_length": self.max_length,160 "rep_pen": self.rep_pen,161 "rep_pen_range": self.rep_pen_range,162 "rep_pen_slope": self.rep_pen_slope,163 "temperature": self.temperature,164 "tfs": self.tfs,165 "top_a": self.top_a,166 "top_p": self.top_p,167 "top_k": self.top_k,168 "typical": self.typical,169 }170 171 if stop is not None:172 data["stop_sequence"] = stop173 174 response = requests.post(175 f"{clean_url(self.endpoint)}/api/v1/generate", json=data176 )177 178 response.raise_for_status()179 json_response = response.json()180 181 if (182 "results" in json_response183 and len(json_response["results"]) > 0184 and "text" in json_response["results"][0]185 ):186 text = json_response["results"][0]["text"].strip()187 188 if stop is not None:189 for sequence in stop:190 if text.endswith(sequence):191 text = text[: -len(sequence)].rstrip()192 193 return text194 else:195 raise ValueError(196 f"Unexpected response format from Kobold API: {json_response}"197 )198 