Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
utils.py727 linesDownload Raw Back to tests
1#!/usr/bin/env python32# -*- coding: utf-8 -*-3 4# type: ignore[reportUnusedImport]5 6import subprocess7import os8 9TMP_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "tmp")10import re11import json12from json import JSONDecodeError13import sys14import requests15import time16from concurrent.futures import ThreadPoolExecutor, as_completed17from typing import (18    Any,19    Callable,20    ContextManager,21    Iterable,22    Iterator,23    List,24    Literal,25    Tuple,26    Set,27)28from re import RegexFlag29import wget30 31 32DEFAULT_HTTP_TIMEOUT = 6033 34# per-request timeout, a hung server fails the test instead of stalling the CI for hours35DEFAULT_REQUEST_TIMEOUT = 60036 37 38class ServerResponse:39    headers: dict40    status_code: int41    body: dict | Any42 43 44class ServerError(Exception):45    def __init__(self, code, body):46        self.code = code47        self.body = body48 49 50class ServerProcess:51    # default options52    debug: bool = False53    server_port: int = 808054    server_host: str = "127.0.0.1"55    model_hf_repo: str | None = "ggml-org/models"56    model_hf_file: str | None = "tinyllamas/stories260K.gguf"57    model_alias: str = "tinyllama-2"58    temperature: float = 0.859    seed: int = 4260    offline: bool = False61 62    # custom options63    model_alias: str | None = None64    model_tags: str | None = None65    model_url: str | None = None66    model_file: str | None = None67    model_draft: str | None = None68    n_threads: int | None = None69    n_gpu_layer: int | None = None70    n_batch: int | None = None71    n_ubatch: int | None = None72    n_ctx: int | None = None73    n_ga: int | None = None74    n_ga_w: int | None = None75    n_predict: int | None = None76    n_prompts: int | None = 077    slot_save_path: str | None = None78    id_slot: int | None = None79    cache_prompt: bool | None = None80    n_slots: int | None = None81    ctk: str | None = None82    ctv: str | None = None83    fa: str | None = None84    server_continuous_batching: bool | None = False85    server_embeddings: bool | None = False86    server_reranking: bool | None = False87    server_metrics: bool | None = False88    kv_unified: bool | None = False89    swa_full: bool | None = False90    server_slots: bool | None = False91    pooling: str | None = None92    api_key: str | None = None93    models_dir: str | None = None94    models_max: int | None = None95    models_preset: str | None = None96    no_models_autoload: bool | None = None97    lora_files: List[str] | None = None98    enable_ctx_shift: int | None = False99    spec_type: str | None = None100    spec_draft_n_min: int | None = None101    spec_draft_n_max: int | None = None102    spec_synth_len: float | None = None103    spec_synth_rates: List[float] | None = None104    no_ui: bool | None = None105    jinja: bool | None = None106    reasoning_format: Literal['deepseek', 'none', 'nothink'] | None = None107    reasoning: Literal['on', 'off', 'auto'] | None = None108    chat_template: str | None = None109    chat_template_file: str | None = None110    server_path: str | None = None111    mmproj_url: str | None = None112    no_mmproj: bool | None = None113    media_path: str | None = None114    sleep_idle_seconds: int | None = None115    cache_ram: int | None = None116    no_cache_idle_slots: bool = False117    log_path: str | None = None118    ui_mcp_proxy: bool = False119    backend_sampling: bool = False120    gcp_compat: bool = False121    server_tools: str | None = None122    server_tools_runtime: str | None = None123    mcp_servers_config: str | None = None124    mcp_servers_json: str | None = None125    cors_origins: str | None = None126 127    # session variables128    process: subprocess.Popen | None = None129 130    def __init__(self):131        if "N_GPU_LAYERS" in os.environ:132            self.n_gpu_layer = int(os.environ["N_GPU_LAYERS"])133        if "DEBUG" in os.environ:134            self.debug = True135        if "PORT" in os.environ:136            self.server_port = int(os.environ["PORT"])137        self.external_server = "DEBUG_EXTERNAL" in os.environ138 139    def start(self, timeout_seconds: int = DEFAULT_HTTP_TIMEOUT) -> None:140        env = {141            **os.environ,142            "LLAMA_SERVER_DEBUG_FAKE_TIMING": "1",143        }144        if "LLAMA_CACHE" not in os.environ:145            env["LLAMA_CACHE"] = "tmp"146        if self.external_server:147            print(f"[external_server]: Assuming external server running on {self.server_host}:{self.server_port}")148            return149        if self.server_path is not None:150            server_path = self.server_path151        elif "LLAMA_SERVER_BIN_PATH" in os.environ:152            server_path = os.environ["LLAMA_SERVER_BIN_PATH"]153        elif os.name == "nt":154            server_path = "../../../build/bin/Release/llama-server.exe"155        else:156            server_path = "../../../build/bin/llama-server"157        server_args = [158            "--host",159            self.server_host,160            "--port",161            self.server_port,162            "--temp",163            self.temperature,164            "--seed",165            self.seed,166        ]167        if self.offline:168            server_args.append("--offline")169        if self.model_file:170            server_args.extend(["--model", self.model_file])171        if self.model_url:172            server_args.extend(["--model-url", self.model_url])173        if self.model_draft:174            server_args.extend(["--model-draft", self.model_draft])175        if self.model_hf_repo:176            server_args.extend(["--hf-repo", self.model_hf_repo])177        if self.model_hf_file:178            server_args.extend(["--hf-file", self.model_hf_file])179        if self.models_dir:180            server_args.extend(["--models-dir", self.models_dir])181        if self.models_max is not None:182            server_args.extend(["--models-max", self.models_max])183        if self.models_preset:184            server_args.extend(["--models-preset", self.models_preset])185        if self.cors_origins:186            server_args.extend(["--cors-origins", self.cors_origins])187        if self.n_batch:188            server_args.extend(["--batch-size", self.n_batch])189        if self.n_ubatch:190            server_args.extend(["--ubatch-size", self.n_ubatch])191        if self.n_threads:192            server_args.extend(["--threads", self.n_threads])193        if self.n_gpu_layer:194            server_args.extend(["--n-gpu-layers", self.n_gpu_layer])195        if self.server_continuous_batching:196            server_args.append("--cont-batching")197        if self.server_embeddings:198            server_args.append("--embedding")199        if self.server_reranking:200            server_args.append("--reranking")201        if self.server_metrics:202            server_args.append("--metrics")203        if self.kv_unified:204            server_args.append("--kv-unified")205        if self.swa_full:206            server_args.append("--swa-full")207        if self.server_slots:208            server_args.append("--slots")209        else:210            server_args.append("--no-slots")211        if self.pooling:212            server_args.extend(["--pooling", self.pooling])213        if self.model_alias:214            server_args.extend(["--alias", self.model_alias])215        if self.model_tags:216            server_args.extend(["--tags", self.model_tags])217        if self.n_ctx:218            server_args.extend(["--ctx-size", self.n_ctx])219        if self.n_slots:220            server_args.extend(["--parallel", self.n_slots])221        if self.ctk:222            server_args.extend(["-ctk", self.ctk])223        if self.ctv:224            server_args.extend(["-ctv", self.ctv])225        if self.fa is not None:226            server_args.extend(["-fa", self.fa])227        if self.n_predict:228            server_args.extend(["--n-predict", self.n_predict])229        if self.slot_save_path:230            server_args.extend(["--slot-save-path", self.slot_save_path])231        if self.n_ga:232            server_args.extend(["--grp-attn-n", self.n_ga])233        if self.n_ga_w:234            server_args.extend(["--grp-attn-w", self.n_ga_w])235        if self.debug:236            server_args.append("--verbose")237        if self.lora_files:238            for lora_file in self.lora_files:239                server_args.extend(["--lora", lora_file])240        if self.enable_ctx_shift:241            server_args.append("--context-shift")242        if self.spec_type:243            server_args.extend(["--spec-type", self.spec_type])244        if self.api_key:245            server_args.extend(["--api-key", self.api_key])246        if self.spec_draft_n_max:247            server_args.extend(["--spec-draft-n-max", self.spec_draft_n_max])248        if self.spec_draft_n_min:249            server_args.extend(["--spec-draft-n-min", self.spec_draft_n_min])250        if self.spec_synth_len is not None:251            server_args.extend(["--spec-synth-len", self.spec_synth_len])252        if self.spec_synth_rates is not None:253            rates = ",".join(str(rate) for rate in self.spec_synth_rates)254            server_args.extend(["--spec-synth-rates", rates])255        if self.no_ui:256            server_args.append("--no-ui")257        if self.no_models_autoload:258            server_args.append("--no-models-autoload")259        if self.jinja:260            server_args.append("--jinja")261        else:262            server_args.append("--no-jinja")263        if self.reasoning_format is not None:264            server_args.extend(("--reasoning-format", self.reasoning_format))265        if self.reasoning is not None:266            server_args.extend(("--reasoning", self.reasoning))267        if self.chat_template:268            server_args.extend(["--chat-template", self.chat_template])269        if self.chat_template_file:270            server_args.extend(["--chat-template-file", self.chat_template_file])271        if self.mmproj_url:272            server_args.extend(["--mmproj-url", self.mmproj_url])273        if self.no_mmproj:274            server_args.append("--no-mmproj")275        if self.media_path:276            server_args.extend(["--media-path", self.media_path])277        if self.sleep_idle_seconds is not None:278            server_args.extend(["--sleep-idle-seconds", self.sleep_idle_seconds])279        if self.cache_ram is not None:280            server_args.extend(["--cache-ram", self.cache_ram])281        if self.no_cache_idle_slots:282            server_args.append("--no-cache-idle-slots")283        if self.ui_mcp_proxy:284            server_args.append("--ui-mcp-proxy")285        if self.server_tools:286            server_args.extend(["--tools", self.server_tools])287        if self.server_tools_runtime:288            server_args.extend(["--tools-runtime", self.server_tools_runtime])289        if self.mcp_servers_config:290            server_args.extend(["--mcp-servers-config", self.mcp_servers_config])291        if self.mcp_servers_json:292            server_args.extend(["--mcp-servers-json", self.mcp_servers_json])293        if self.backend_sampling:294            server_args.append("--backend_sampling")295        if self.gcp_compat:296            env["AIP_MODE"] = "PREDICTION"297            env["AIP_HTTP_PORT"] = str(self.server_port)298 299        args = [str(arg) for arg in [server_path, *server_args]]300        print(f"tests: starting server with: {' '.join(args)}")301 302        flags = 0303        if "nt" == os.name:304            flags |= subprocess.DETACHED_PROCESS305            flags |= subprocess.CREATE_NEW_PROCESS_GROUP306            flags |= subprocess.CREATE_NO_WINDOW307 308        if self.log_path:309            self._log = open(self.log_path, "w")310        else:311            self._log = sys.stdout312 313        self.process = subprocess.Popen(314            [str(arg) for arg in [server_path, *server_args]],315            creationflags=flags,316            stdout=self._log,317            stderr=self._log if self._log != sys.stdout else sys.stdout,318            env=env,319        )320        server_instances.add(self)321 322        print(f"server pid={self.process.pid}, pytest pid={os.getpid()}")323 324        # wait for server to start325        start_time = time.time()326        last_print_time = start_time327        while time.time() - start_time < timeout_seconds:328            try:329                response = self.make_request("GET", "/health", headers={330                    "Authorization": f"Bearer {self.api_key}" if self.api_key else None331                })332                if response.status_code == 200:333                    self.ready = True334                    return  # server is ready335            except Exception as e:336                pass337            # Check if process died338            if self.process.poll() is not None:339                raise RuntimeError(f"Server process died with return code {self.process.returncode}")340 341            if time.time() - last_print_time >= 1.0:342                print(f"Waiting for server to start...")343                last_print_time = time.time()344            time.sleep(0.01)345        raise TimeoutError(f"Server did not start within {timeout_seconds} seconds")346 347    def stop(self) -> None:348        if self.external_server:349            print("[external_server]: Not stopping external server")350            return351        if self in server_instances:352            server_instances.remove(self)353        if self.process:354            print(f"Stopping server with pid={self.process.pid}")355            self.process.terminate()356            try:357                self.process.wait(timeout=5)358            except subprocess.TimeoutExpired:359                print(f"Server pid={self.process.pid} did not terminate in time, killing")360                self.process.kill()361                self.process.wait(timeout=5)362            except Exception as e:363                print(f"Error waiting for server: {e}")364            self.process = None365        if hasattr(self, '_log') and self._log != sys.stdout:366            self._log.close()367 368    def make_request(369        self,370        method: str,371        path: str,372        data: dict | Any | None = None,373        headers: dict | None = None,374        timeout: float | None = DEFAULT_REQUEST_TIMEOUT,375    ) -> ServerResponse:376        url = f"http://{self.server_host}:{self.server_port}{path}"377        parse_body = False378        if method == "GET":379            response = requests.get(url, headers=headers, timeout=timeout)380            parse_body = True381        elif method == "POST":382            response = requests.post(url, headers=headers, json=data, timeout=timeout)383            parse_body = True384        elif method == "DELETE":385            response = requests.delete(url, headers=headers, timeout=timeout)386            parse_body = True387        elif method == "OPTIONS":388            response = requests.options(url, headers=headers, timeout=timeout)389        else:390            raise ValueError(f"Unimplemented method: {method}")391        result = ServerResponse()392        result.headers = dict(response.headers)393        result.status_code = response.status_code394        if parse_body:395            try:396                result.body = response.json()397            except (JSONDecodeError, requests.exceptions.JSONDecodeError):398                result.body = response.text399        else:400            result.body = None401        print("Response from server", json.dumps(result.body, indent=2))402        return result403 404    def make_stream_request(405        self,406        method: str,407        path: str,408        data: dict | None = None,409        headers: dict | None = None,410    ) -> Iterator[dict]:411        url = f"http://{self.server_host}:{self.server_port}{path}"412        if method == "POST":413            response = requests.post(url, headers=headers, json=data, stream=True)414        else:415            raise ValueError(f"Unimplemented method: {method}")416        if response.status_code != 200:417            raise ServerError(response.status_code, response.json())418        for line_bytes in response.iter_lines():419            line = line_bytes.decode("utf-8")420            if '[DONE]' in line:421                break422            elif line.startswith('data: '):423                data = json.loads(line[6:])424                print("Partial response from server", json.dumps(data, indent=2))425                yield data426 427    def make_any_request(428        self,429        method: str,430        path: str,431        data: dict | None = None,432        headers: dict | None = None,433        timeout: float | None = DEFAULT_REQUEST_TIMEOUT,434    ) -> dict:435        stream = data.get('stream', False)436        if stream:437            content: list[str] = []438            reasoning_content: list[str] = []439            tool_calls: list[dict] = []440            finish_reason: Optional[str] = None441 442            content_parts = 0443            reasoning_content_parts = 0444            tool_call_parts = 0445            arguments_parts = 0446 447            for chunk in self.make_stream_request(method, path, data, headers):448                if chunk['choices']:449                    assert len(chunk['choices']) == 1, f'Expected 1 choice, got {len(chunk["choices"])}'450                    choice = chunk['choices'][0]451                    if choice['delta'].get('content') is not None:452                        assert len(choice['delta']['content']) > 0, f'Expected non empty content delta!'453                        content.append(choice['delta']['content'])454                        content_parts += 1455                    if choice['delta'].get('reasoning_content') is not None:456                        assert len(choice['delta']['reasoning_content']) > 0, f'Expected non empty reasoning_content delta!'457                        reasoning_content.append(choice['delta']['reasoning_content'])458                        reasoning_content_parts += 1459                    if choice['delta'].get('finish_reason') is not None:460                        finish_reason = choice['delta']['finish_reason']461                    for tc in choice['delta'].get('tool_calls', []):462                        if 'function' not in tc:463                            raise ValueError(f"Expected function type, got {tc['type']}")464                        if tc['index'] >= len(tool_calls):465                            assert 'id' in tc466                            assert tc.get('type') == 'function'467                            assert 'function' in tc and 'name' in tc['function'] and len(tc['function']['name']) > 0, \468                                f"Expected function call with name, got {tc.get('function')}"469                            tool_calls.append(dict(470                                id="",471                                type="function",472                                function=dict(473                                    name="",474                                    arguments="",475                                )476                            ))477                        tool_call = tool_calls[tc['index']]478                        if tc.get('id') is not None:479                            tool_call['id'] = tc['id']480                        fct = tc['function']481                        assert 'id' not in fct, f"Function call should not have id: {fct}"482                        if fct.get('name') is not None:483                            tool_call['function']['name'] = tool_call['function'].get('name', '') + fct['name']484                        if fct.get('arguments') is not None:485                            tool_call['function']['arguments'] += fct['arguments']486                            arguments_parts += 1487                        tool_call_parts += 1488                else:489                    # When `include_usage` is True (the default), we expect the last chunk of the stream490                    # immediately preceding the `data: [DONE]` message to contain a `choices` field with an empty array491                    # and a `usage` field containing the usage statistics (n.b., llama-server also returns `timings` in492                    # the last chunk)493                    assert 'usage' in chunk, f"Expected finish_reason in chunk: {chunk}"494                    assert 'timings' in chunk, f"Expected finish_reason in chunk: {chunk}"495            print(f'Streamed response had {content_parts} content parts, {reasoning_content_parts} reasoning_content parts, {tool_call_parts} tool call parts incl. {arguments_parts} arguments parts')496            result = dict(497                choices=[498                    dict(499                        index=0,500                        finish_reason=finish_reason,501                        message=dict(502                            role='assistant',503                            content=''.join(content) if content else None,504                            reasoning_content=''.join(reasoning_content) if reasoning_content else None,505                            tool_calls=tool_calls if tool_calls else None,506                        ),507                    )508                ],509            )510            print("Final response from server", json.dumps(result, indent=2))511            return result512        else:513            response = self.make_request(method, path, data, headers, timeout=timeout)514            assert response.status_code == 200, f"Server returned error: {response.status_code}"515            return response.body516 517 518 519server_instances: Set[ServerProcess] = set()520 521 522class ServerPreset:523    @staticmethod524    def load_all() -> None:525        """ Load all server presets to ensure model files are cached. """526        servers: List[ServerProcess] = [527            method()528            for name, method in ServerPreset.__dict__.items()529            if callable(method) and name != "load_all"530        ]531        for server in servers:532            server.offline = False533            server.start()534            server.stop()535 536    @staticmethod537    def tinyllama2() -> ServerProcess:538        server = ServerProcess()539        server.offline = True # will be downloaded by load_all()540        server.model_hf_repo = "ggml-org/test-model-stories260K"541        server.model_hf_file = None542        server.model_alias = "tinyllama-2"543        server.n_ctx = 512544        server.n_batch = 32545        server.n_slots = 2546        server.n_predict = 64547        server.seed = 42548        return server549 550    @staticmethod551    def bert_bge_small() -> ServerProcess:552        server = ServerProcess()553        server.offline = True # will be downloaded by load_all()554        server.model_hf_repo = "ggml-org/models"555        server.model_hf_file = "bert-bge-small/ggml-model-f16.gguf"556        server.model_alias = "bert-bge-small"557        server.n_ctx = 512558        server.n_batch = 128559        server.n_ubatch = 128560        server.n_slots = 2561        server.seed = 42562        server.server_embeddings = True563        return server564 565    @staticmethod566    def bert_bge_small_with_fa() -> ServerProcess:567        server = ServerProcess()568        server.offline = True # will be downloaded by load_all()569        server.model_hf_repo = "ggml-org/models"570        server.model_hf_file = "bert-bge-small/ggml-model-f16.gguf"571        server.model_alias = "bert-bge-small"572        server.n_ctx = 1024573        server.n_batch = 300574        server.n_ubatch = 300575        server.n_slots = 2576        server.fa = "on"577        server.seed = 42578        server.server_embeddings = True579        return server580 581    @staticmethod582    def tinyllama_infill() -> ServerProcess:583        server = ServerProcess()584        server.offline = True # will be downloaded by load_all()585        server.model_hf_repo = "ggml-org/test-model-stories260K-infill"586        server.model_hf_file = None587        server.model_alias = "tinyllama-infill"588        server.n_ctx = 2048589        server.n_batch = 1024590        server.n_slots = 1591        server.n_predict = 64592        server.temperature = 0.0593        server.seed = 42594        return server595 596    @staticmethod597    def stories15m_moe() -> ServerProcess:598        server = ServerProcess()599        server.offline = True # will be downloaded by load_all()600        server.model_hf_repo = "ggml-org/stories15M_MOE"601        server.model_hf_file = "stories15M_MOE-F16.gguf"602        server.model_alias = "stories15m-moe"603        server.n_ctx = 2048604        server.n_batch = 1024605        server.n_slots = 1606        server.n_predict = 64607        server.temperature = 0.0608        server.seed = 42609        return server610 611    @staticmethod612    def jina_reranker_tiny() -> ServerProcess:613        server = ServerProcess()614        server.offline = True # will be downloaded by load_all()615        server.model_hf_repo = "ggml-org/models"616        server.model_hf_file = "jina-reranker-v1-tiny-en/ggml-model-f16.gguf"617        server.model_alias = "jina-reranker"618        server.n_ctx = 512619        server.n_batch = 512620        server.n_slots = 1621        server.seed = 42622        server.server_reranking = True623        return server624 625    @staticmethod626    def tinygemma3() -> ServerProcess:627        server = ServerProcess()628        server.offline = True # will be downloaded by load_all()629        # mmproj is already provided by HF registry API630        server.model_hf_file = None631        server.model_hf_repo = "ggml-org/tinygemma3-GGUF:Q8_0"632        server.model_alias = "tinygemma3"633        server.n_ctx = 1024634        server.n_batch = 512635        server.n_slots = 2636        server.n_predict = 4637        server.seed = 42638        return server639 640    @staticmethod641    def router() -> ServerProcess:642        server = ServerProcess()643        server.offline = True # will be downloaded by load_all()644        # router server has no models645        server.model_file = None646        server.model_alias = None647        server.model_hf_repo = None648        server.model_hf_file = None649        server.n_ctx = 1024650        server.n_batch = 16651        server.n_slots = 1652        server.n_predict = 16653        server.seed = 42654        return server655 656 657def parallel_function_calls(function_list: List[Tuple[Callable[..., Any], Tuple[Any, ...]]]) -> List[Any]:658    """659    Run multiple functions in parallel and return results in the same order as calls. Equivalent to Promise.all in JS.660 661    Example usage:662 663    results = parallel_function_calls([664        (func1, (arg1, arg2)),665        (func2, (arg3, arg4)),666    ])667    """668    results = [None] * len(function_list)669    exceptions = []670 671    def worker(index, func, args):672        try:673            result = func(*args)674            results[index] = result675        except Exception as e:676            exceptions.append((index, str(e)))677 678    with ThreadPoolExecutor() as executor:679        futures = []680        for i, (func, args) in enumerate(function_list):681            future = executor.submit(worker, i, func, args)682            futures.append(future)683 684        # Wait for all futures to complete685        for future in as_completed(futures):686            pass687 688    # Check if there were any exceptions689    if exceptions:690        print("Exceptions occurred:")691        for index, error in exceptions:692            print(f"Function at index {index}: {error}")693 694    return results695 696 697def match_regex(regex: str, text: str) -> bool:698    return (699        re.compile(700            regex, flags=RegexFlag.IGNORECASE | RegexFlag.MULTILINE | RegexFlag.DOTALL701        ).search(text)702        is not None703    )704 705 706def download_file(url: str, output_file_path: str | None = None) -> str:707    """708    Download a file from a URL to a local path. If the file already exists, it will not be downloaded again.709 710    output_file_path is the local path to save the downloaded file. If not provided, the file will be saved in the root directory.711 712    Returns the local path of the downloaded file.713    """714    file_name = url.split('/').pop()715    output_file = f'./tmp/{file_name}' if output_file_path is None else output_file_path716    if not os.path.exists(output_file):717        print(f"Downloading {url} to {output_file}")718        wget.download(url, out=output_file)719        print(f"Done downloading to {output_file}")720    else:721        print(f"File already exists at {output_file}")722    return output_file723 724 725def is_slow_test_allowed():726    return os.environ.get("SLOW_TESTS") == "1" or os.environ.get("SLOW_TESTS") == "ON"727