KBaba7/llama.cpp
0
1#!/usr/bin/env python32# -*- coding: utf-8 -*-3 4# type: ignore[reportUnusedImport]5 6import subprocess7import os8import re9import json10import sys11import requests12import time13from concurrent.futures import ThreadPoolExecutor, as_completed14from typing import (15 Any,16 Callable,17 ContextManager,18 Iterable,19 Iterator,20 List,21 Literal,22 Tuple,23 Set,24)25from re import RegexFlag26import wget27 28 29DEFAULT_HTTP_TIMEOUT = 12 if "LLAMA_SANITIZE" not in os.environ else 3030 31 32class ServerResponse:33 headers: dict34 status_code: int35 body: dict | Any36 37 38class ServerProcess:39 # default options40 debug: bool = False41 server_port: int = 808042 server_host: str = "127.0.0.1"43 model_hf_repo: str = "ggml-org/models"44 model_hf_file: str | None = "tinyllamas/stories260K.gguf"45 model_alias: str = "tinyllama-2"46 temperature: float = 0.847 seed: int = 4248 49 # custom options50 model_alias: str | None = None51 model_url: str | None = None52 model_file: str | None = None53 model_draft: str | None = None54 n_threads: int | None = None55 n_gpu_layer: int | None = None56 n_batch: int | None = None57 n_ubatch: int | None = None58 n_ctx: int | None = None59 n_ga: int | None = None60 n_ga_w: int | None = None61 n_predict: int | None = None62 n_prompts: int | None = 063 slot_save_path: str | None = None64 id_slot: int | None = None65 cache_prompt: bool | None = None66 n_slots: int | None = None67 server_continuous_batching: bool | None = False68 server_embeddings: bool | None = False69 server_reranking: bool | None = False70 server_metrics: bool | None = False71 server_slots: bool | None = False72 pooling: str | None = None73 draft: int | None = None74 api_key: str | None = None75 lora_files: List[str] | None = None76 disable_ctx_shift: int | None = False77 draft_min: int | None = None78 draft_max: int | None = None79 no_webui: bool | None = None80 jinja: bool | None = None81 chat_template: str | None = None82 chat_template_file: str | None = None83 84 # session variables85 process: subprocess.Popen | None = None86 87 def __init__(self):88 if "N_GPU_LAYERS" in os.environ:89 self.n_gpu_layer = int(os.environ["N_GPU_LAYERS"])90 if "DEBUG" in os.environ:91 self.debug = True92 if "PORT" in os.environ:93 self.server_port = int(os.environ["PORT"])94 95 def start(self, timeout_seconds: int | None = DEFAULT_HTTP_TIMEOUT) -> None:96 if "LLAMA_SERVER_BIN_PATH" in os.environ:97 server_path = os.environ["LLAMA_SERVER_BIN_PATH"]98 elif os.name == "nt":99 server_path = "../../../build/bin/Release/llama-server.exe"100 else:101 server_path = "../../../build/bin/llama-server"102 server_args = [103 "--host",104 self.server_host,105 "--port",106 self.server_port,107 "--temp",108 self.temperature,109 "--seed",110 self.seed,111 ]112 if self.model_file:113 server_args.extend(["--model", self.model_file])114 if self.model_url:115 server_args.extend(["--model-url", self.model_url])116 if self.model_draft:117 server_args.extend(["--model-draft", self.model_draft])118 if self.model_hf_repo:119 server_args.extend(["--hf-repo", self.model_hf_repo])120 if self.model_hf_file:121 server_args.extend(["--hf-file", self.model_hf_file])122 if self.n_batch:123 server_args.extend(["--batch-size", self.n_batch])124 if self.n_ubatch:125 server_args.extend(["--ubatch-size", self.n_ubatch])126 if self.n_threads:127 server_args.extend(["--threads", self.n_threads])128 if self.n_gpu_layer:129 server_args.extend(["--n-gpu-layers", self.n_gpu_layer])130 if self.draft is not None:131 server_args.extend(["--draft", self.draft])132 if self.server_continuous_batching:133 server_args.append("--cont-batching")134 if self.server_embeddings:135 server_args.append("--embedding")136 if self.server_reranking:137 server_args.append("--reranking")138 if self.server_metrics:139 server_args.append("--metrics")140 if self.server_slots:141 server_args.append("--slots")142 if self.pooling:143 server_args.extend(["--pooling", self.pooling])144 if self.model_alias:145 server_args.extend(["--alias", self.model_alias])146 if self.n_ctx:147 server_args.extend(["--ctx-size", self.n_ctx])148 if self.n_slots:149 server_args.extend(["--parallel", self.n_slots])150 if self.n_predict:151 server_args.extend(["--n-predict", self.n_predict])152 if self.slot_save_path:153 server_args.extend(["--slot-save-path", self.slot_save_path])154 if self.n_ga:155 server_args.extend(["--grp-attn-n", self.n_ga])156 if self.n_ga_w:157 server_args.extend(["--grp-attn-w", self.n_ga_w])158 if self.debug:159 server_args.append("--verbose")160 if self.lora_files:161 for lora_file in self.lora_files:162 server_args.extend(["--lora", lora_file])163 if self.disable_ctx_shift:164 server_args.extend(["--no-context-shift"])165 if self.api_key:166 server_args.extend(["--api-key", self.api_key])167 if self.draft_max:168 server_args.extend(["--draft-max", self.draft_max])169 if self.draft_min:170 server_args.extend(["--draft-min", self.draft_min])171 if self.no_webui:172 server_args.append("--no-webui")173 if self.jinja:174 server_args.append("--jinja")175 if self.chat_template:176 server_args.extend(["--chat-template", self.chat_template])177 if self.chat_template_file:178 server_args.extend(["--chat-template-file", self.chat_template_file])179 180 args = [str(arg) for arg in [server_path, *server_args]]181 print(f"bench: starting server with: {' '.join(args)}")182 183 flags = 0184 if "nt" == os.name:185 flags |= subprocess.DETACHED_PROCESS186 flags |= subprocess.CREATE_NEW_PROCESS_GROUP187 flags |= subprocess.CREATE_NO_WINDOW188 189 self.process = subprocess.Popen(190 [str(arg) for arg in [server_path, *server_args]],191 creationflags=flags,192 stdout=sys.stdout,193 stderr=sys.stdout,194 env={**os.environ, "LLAMA_CACHE": "tmp"} if "LLAMA_CACHE" not in os.environ else None,195 )196 server_instances.add(self)197 198 print(f"server pid={self.process.pid}, pytest pid={os.getpid()}")199 200 # wait for server to start201 start_time = time.time()202 while time.time() - start_time < timeout_seconds:203 try:204 response = self.make_request("GET", "/health", headers={205 "Authorization": f"Bearer {self.api_key}" if self.api_key else None206 })207 if response.status_code == 200:208 self.ready = True209 return # server is ready210 except Exception as e:211 pass212 print(f"Waiting for server to start...")213 time.sleep(0.5)214 raise TimeoutError(f"Server did not start within {timeout_seconds} seconds")215 216 def stop(self) -> None:217 if self in server_instances:218 server_instances.remove(self)219 if self.process:220 print(f"Stopping server with pid={self.process.pid}")221 self.process.kill()222 self.process = None223 224 def make_request(225 self,226 method: str,227 path: str,228 data: dict | Any | None = None,229 headers: dict | None = None,230 timeout: float | None = None,231 ) -> ServerResponse:232 url = f"http://{self.server_host}:{self.server_port}{path}"233 parse_body = False234 if method == "GET":235 response = requests.get(url, headers=headers, timeout=timeout)236 parse_body = True237 elif method == "POST":238 response = requests.post(url, headers=headers, json=data, timeout=timeout)239 parse_body = True240 elif method == "OPTIONS":241 response = requests.options(url, headers=headers, timeout=timeout)242 else:243 raise ValueError(f"Unimplemented method: {method}")244 result = ServerResponse()245 result.headers = dict(response.headers)246 result.status_code = response.status_code247 result.body = response.json() if parse_body else None248 print("Response from server", json.dumps(result.body, indent=2))249 return result250 251 def make_stream_request(252 self,253 method: str,254 path: str,255 data: dict | None = None,256 headers: dict | None = None,257 ) -> Iterator[dict]:258 url = f"http://{self.server_host}:{self.server_port}{path}"259 if method == "POST":260 response = requests.post(url, headers=headers, json=data, stream=True)261 else:262 raise ValueError(f"Unimplemented method: {method}")263 for line_bytes in response.iter_lines():264 line = line_bytes.decode("utf-8")265 if '[DONE]' in line:266 break267 elif line.startswith('data: '):268 data = json.loads(line[6:])269 print("Partial response from server", json.dumps(data, indent=2))270 yield data271 272 273server_instances: Set[ServerProcess] = set()274 275 276class ServerPreset:277 @staticmethod278 def tinyllama2() -> ServerProcess:279 server = ServerProcess()280 server.model_hf_repo = "ggml-org/models"281 server.model_hf_file = "tinyllamas/stories260K.gguf"282 server.model_alias = "tinyllama-2"283 server.n_ctx = 256284 server.n_batch = 32285 server.n_slots = 2286 server.n_predict = 64287 server.seed = 42288 return server289 290 @staticmethod291 def bert_bge_small() -> ServerProcess:292 server = ServerProcess()293 server.model_hf_repo = "ggml-org/models"294 server.model_hf_file = "bert-bge-small/ggml-model-f16.gguf"295 server.model_alias = "bert-bge-small"296 server.n_ctx = 512297 server.n_batch = 128298 server.n_ubatch = 128299 server.n_slots = 2300 server.seed = 42301 server.server_embeddings = True302 return server303 304 @staticmethod305 def tinyllama_infill() -> ServerProcess:306 server = ServerProcess()307 server.model_hf_repo = "ggml-org/models"308 server.model_hf_file = "tinyllamas/stories260K-infill.gguf"309 server.model_alias = "tinyllama-infill"310 server.n_ctx = 2048311 server.n_batch = 1024312 server.n_slots = 1313 server.n_predict = 64314 server.temperature = 0.0315 server.seed = 42316 return server317 318 @staticmethod319 def stories15m_moe() -> ServerProcess:320 server = ServerProcess()321 server.model_hf_repo = "ggml-org/stories15M_MOE"322 server.model_hf_file = "stories15M_MOE-F16.gguf"323 server.model_alias = "stories15m-moe"324 server.n_ctx = 2048325 server.n_batch = 1024326 server.n_slots = 1327 server.n_predict = 64328 server.temperature = 0.0329 server.seed = 42330 return server331 332 @staticmethod333 def jina_reranker_tiny() -> ServerProcess:334 server = ServerProcess()335 server.model_hf_repo = "ggml-org/models"336 server.model_hf_file = "jina-reranker-v1-tiny-en/ggml-model-f16.gguf"337 server.model_alias = "jina-reranker"338 server.n_ctx = 512339 server.n_batch = 512340 server.n_slots = 1341 server.seed = 42342 server.server_reranking = True343 return server344 345 346def parallel_function_calls(function_list: List[Tuple[Callable[..., Any], Tuple[Any, ...]]]) -> List[Any]:347 """348 Run multiple functions in parallel and return results in the same order as calls. Equivalent to Promise.all in JS.349 350 Example usage:351 352 results = parallel_function_calls([353 (func1, (arg1, arg2)),354 (func2, (arg3, arg4)),355 ])356 """357 results = [None] * len(function_list)358 exceptions = []359 360 def worker(index, func, args):361 try:362 result = func(*args)363 results[index] = result364 except Exception as e:365 exceptions.append((index, str(e)))366 367 with ThreadPoolExecutor() as executor:368 futures = []369 for i, (func, args) in enumerate(function_list):370 future = executor.submit(worker, i, func, args)371 futures.append(future)372 373 # Wait for all futures to complete374 for future in as_completed(futures):375 pass376 377 # Check if there were any exceptions378 if exceptions:379 print("Exceptions occurred:")380 for index, error in exceptions:381 print(f"Function at index {index}: {error}")382 383 return results384 385 386def match_regex(regex: str, text: str) -> bool:387 return (388 re.compile(389 regex, flags=RegexFlag.IGNORECASE | RegexFlag.MULTILINE | RegexFlag.DOTALL390 ).search(text)391 is not None392 )393 394 395def download_file(url: str, output_file_path: str | None = None) -> str:396 """397 Download a file from a URL to a local path. If the file already exists, it will not be downloaded again.398 399 output_file_path is the local path to save the downloaded file. If not provided, the file will be saved in the root directory.400 401 Returns the local path of the downloaded file.402 """403 file_name = url.split('/').pop()404 output_file = f'./tmp/{file_name}' if output_file_path is None else output_file_path405 if not os.path.exists(output_file):406 print(f"Downloading {url} to {output_file}")407 wget.download(url, out=output_file)408 print(f"Done downloading to {output_file}")409 else:410 print(f"File already exists at {output_file}")411 return output_file412 413 414def is_slow_test_allowed():415 return os.environ.get("SLOW_TESTS") == "1" or os.environ.get("SLOW_TESTS") == "ON"416 