Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
utils.py416 linesDownload Raw Back to tests
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