Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
bench.py326 linesDownload Raw Back to bench
1from __future__ import annotations2 3import argparse4import json5import os6import re7import signal8import socket9import subprocess10import sys11import threading12import time13import traceback14from contextlib import closing15from datetime import datetime16 17import matplotlib18import matplotlib.dates19import matplotlib.pyplot as plt20import requests21from statistics import mean22 23 24def main(args_in: list[str] | None = None) -> None:25    parser = argparse.ArgumentParser(description="Start server benchmark scenario")26    parser.add_argument("--name", type=str, help="Bench name", required=True)27    parser.add_argument("--runner-label", type=str, help="Runner label", required=True)28    parser.add_argument("--branch", type=str, help="Branch name", default="detached")29    parser.add_argument("--commit", type=str, help="Commit name", default="dirty")30    parser.add_argument("--host", type=str, help="Server listen host", default="0.0.0.0")31    parser.add_argument("--port", type=int, help="Server listen host", default="8080")32    parser.add_argument("--model-path-prefix", type=str, help="Prefix where to store the model files", default="models")33    parser.add_argument("--n-prompts", type=int,34                        help="SERVER_BENCH_N_PROMPTS: total prompts to randomly select in the benchmark", required=True)35    parser.add_argument("--max-prompt-tokens", type=int,36                        help="SERVER_BENCH_MAX_PROMPT_TOKENS: maximum prompt tokens to filter out in the dataset",37                        required=True)38    parser.add_argument("--max-tokens", type=int,39                        help="SERVER_BENCH_MAX_CONTEXT: maximum context size of the completions request to filter out in the dataset: prompt + predicted tokens",40                        required=True)41    parser.add_argument("--hf-repo", type=str, help="Hugging Face model repository", required=True)42    parser.add_argument("--hf-file", type=str, help="Hugging Face model file", required=True)43    parser.add_argument("--offline", action="store_true", default=False, help="Offline mode: forces use of cache, prevents network access")44    parser.add_argument("-ngl", "--n-gpu-layers", type=int, help="layers to the GPU for computation", required=True)45    parser.add_argument("--ctx-size", type=int, help="Set the size of the prompt context", required=True)46    parser.add_argument("--parallel", type=int, help="Set the number of slots for process requests", required=True)47    parser.add_argument("--batch-size", type=int, help="Set the batch size for prompt processing", required=True)48    parser.add_argument("--ubatch-size", type=int, help="physical maximum batch size", required=True)49    parser.add_argument("--scenario", type=str, help="Scenario to run", required=True)50    parser.add_argument("--duration", type=str, help="Bench scenario", required=True)51 52    args = parser.parse_args(args_in)53 54    start_time = time.time()55 56    # Start the server and performance scenario57    try:58        server_process = start_server(args)59    except Exception:60        print("bench: server start error :")61        traceback.print_exc(file=sys.stdout)62        sys.exit(1)63 64    # start the benchmark65    iterations = 066    data = {}67    try:68        start_benchmark(args)69 70        with open("results.github.env", 'w') as github_env:71            # parse output72            with open('k6-results.json', 'r') as bench_results:73                # Load JSON data from file74                data = json.load(bench_results)75                for metric_name in data['metrics']:76                    for metric_metric in data['metrics'][metric_name]:77                        value = data['metrics'][metric_name][metric_metric]78                        if isinstance(value, float) or isinstance(value, int):79                            value = round(value, 2)80                            data['metrics'][metric_name][metric_metric]=value81                            github_env.write(82                                f"{escape_metric_name(metric_name)}_{escape_metric_name(metric_metric)}={value}\n")83                iterations = data['root_group']['checks']['success completion']['passes']84 85    except Exception:86        print("bench: error :")87        traceback.print_exc(file=sys.stdout)88 89    # Stop the server90    if server_process:91        try:92            print(f"bench: shutting down server pid={server_process.pid} ...")93            if os.name == 'nt':94                interrupt = signal.CTRL_C_EVENT95            else:96                interrupt = signal.SIGINT97            server_process.send_signal(interrupt)98            server_process.wait(0.5)99 100        except subprocess.TimeoutExpired:101            print(f"server still alive after 500ms, force-killing pid={server_process.pid} ...")102            server_process.kill()  # SIGKILL103            server_process.wait()104 105        while is_server_listening(args.host, args.port):106            time.sleep(0.1)107 108    title = (f"llama.cpp {args.name} on {args.runner_label}\n "109             f"duration={args.duration} {iterations} iterations")110    xlabel = (f"{args.hf_repo}/{args.hf_file}\n"111              f"parallel={args.parallel} ctx-size={args.ctx_size} ngl={args.n_gpu_layers} batch-size={args.batch_size} ubatch-size={args.ubatch_size} pp={args.max_prompt_tokens} pp+tg={args.max_tokens}\n"112              f"branch={args.branch} commit={args.commit}")113 114    # Prometheus115    end_time = time.time()116    prometheus_metrics = {}117    if is_server_listening("0.0.0.0", 9090):118        metrics = ['prompt_tokens_seconds', 'predicted_tokens_seconds',119                   'kv_cache_usage_ratio', 'requests_processing', 'requests_deferred']120 121        for metric in metrics:122            resp = requests.get(f"http://localhost:9090/api/v1/query_range",123                                params={'query': 'llamacpp:' + metric, 'start': start_time, 'end': end_time, 'step': 2})124 125            with open(f"{metric}.json", 'w') as metric_json:126                metric_json.write(resp.text)127 128            if resp.status_code != 200:129                print(f"bench: unable to extract prometheus metric {metric}: {resp.text}")130            else:131                metric_data = resp.json()132                values = metric_data['data']['result'][0]['values']133                timestamps, metric_values = zip(*values)134                metric_values = [float(value) for value in metric_values]135                prometheus_metrics[metric] = metric_values136                timestamps_dt = [str(datetime.fromtimestamp(int(ts))) for ts in timestamps]137                plt.figure(figsize=(16, 10), dpi=80)138                plt.plot(timestamps_dt, metric_values, label=metric)139                plt.xticks(rotation=0, fontsize=14, horizontalalignment='center', alpha=.7)140                plt.yticks(fontsize=12, alpha=.7)141 142                ylabel = f"llamacpp:{metric}"143                plt.title(title,144                          fontsize=14, wrap=True)145                plt.grid(axis='both', alpha=.3)146                plt.ylabel(ylabel, fontsize=22)147                plt.xlabel(xlabel, fontsize=14, wrap=True)148                plt.gca().xaxis.set_major_locator(matplotlib.dates.MinuteLocator())149                plt.gca().xaxis.set_major_formatter(matplotlib.dates.DateFormatter("%Y-%m-%d %H:%M:%S"))150                plt.gcf().autofmt_xdate()151 152                # Remove borders153                plt.gca().spines["top"].set_alpha(0.0)154                plt.gca().spines["bottom"].set_alpha(0.3)155                plt.gca().spines["right"].set_alpha(0.0)156                plt.gca().spines["left"].set_alpha(0.3)157 158                # Save the plot as a jpg image159                plt.savefig(f'{metric}.jpg', dpi=60)160                plt.close()161 162                # Mermaid format in case images upload failed163                with open(f"{metric}.mermaid", 'w') as mermaid_f:164                    mermaid = (165                    f"""---166config:167    xyChart:168        titleFontSize: 12169        width: 900170        height: 600171    themeVariables:172        xyChart:173            titleColor: "#000000"174---175xychart-beta176    title "{title}"177    y-axis "llamacpp:{metric}"178    x-axis "llamacpp:{metric}" {int(min(timestamps))} --> {int(max(timestamps))}179    line [{', '.join([str(round(float(value), 2)) for value in metric_values])}]180                    """)181                    mermaid_f.write(mermaid)182 183    # 140 chars max for commit status description184    bench_results = {185        "i": iterations,186        "req": {187            "p95": round(data['metrics']["http_req_duration"]["p(95)"], 2),188            "avg": round(data['metrics']["http_req_duration"]["avg"], 2),189        },190        "pp": {191            "p95": round(data['metrics']["llamacpp_prompt_processing_second"]["p(95)"], 2),192            "avg": round(data['metrics']["llamacpp_prompt_processing_second"]["avg"], 2),193            "0": round(mean(prometheus_metrics['prompt_tokens_seconds']), 2) if 'prompt_tokens_seconds' in prometheus_metrics else 0,194        },195        "tg": {196            "p95": round(data['metrics']["llamacpp_tokens_second"]["p(95)"], 2),197            "avg": round(data['metrics']["llamacpp_tokens_second"]["avg"], 2),198            "0": round(mean(prometheus_metrics['predicted_tokens_seconds']), 2) if 'predicted_tokens_seconds' in prometheus_metrics else 0,199        },200    }201    with open("results.github.env", 'a') as github_env:202        github_env.write(f"BENCH_RESULTS={json.dumps(bench_results, indent=None, separators=(',', ':') )}\n")203        github_env.write(f"BENCH_ITERATIONS={iterations}\n")204 205        title = title.replace('\n', ' ')206        xlabel = xlabel.replace('\n', ' ')207        github_env.write(f"BENCH_GRAPH_TITLE={title}\n")208        github_env.write(f"BENCH_GRAPH_XLABEL={xlabel}\n")209 210 211def start_benchmark(args):212    k6_path = './k6'213    if 'BENCH_K6_BIN_PATH' in os.environ:214        k6_path = os.environ['BENCH_K6_BIN_PATH']215    k6_args = [216        'run', args.scenario,217        '--no-color',218        '--no-connection-reuse',219        '--no-vu-connection-reuse',220    ]221    k6_args.extend(['--duration', args.duration])222    k6_args.extend(['--iterations', args.n_prompts])223    k6_args.extend(['--vus', args.parallel])224    k6_args.extend(['--summary-export', 'k6-results.json'])225    k6_args.extend(['--out', 'csv=k6-results.csv'])226    args = f"SERVER_BENCH_N_PROMPTS={args.n_prompts} SERVER_BENCH_MAX_PROMPT_TOKENS={args.max_prompt_tokens} SERVER_BENCH_MAX_CONTEXT={args.max_tokens} "227    args = args + ' '.join([str(arg) for arg in [k6_path, *k6_args]])228    print(f"bench: starting k6 with: {args}")229    k6_completed = subprocess.run(args, shell=True, stdout=sys.stdout, stderr=sys.stderr)230    if k6_completed.returncode != 0:231        raise Exception("bench: unable to run k6")232 233 234def start_server(args):235    server_process = start_server_background(args)236 237    attempts = 0238    max_attempts = 600239    if 'GITHUB_ACTIONS' in os.environ:240        max_attempts *= 2241 242    while not is_server_listening(args.host, args.port):243        attempts += 1244        if attempts > max_attempts:245            assert False, "server not started"246        print(f"bench:     waiting for server to start ...")247        time.sleep(0.5)248 249    attempts = 0250    while not is_server_ready(args.host, args.port):251        attempts += 1252        if attempts > max_attempts:253            assert False, "server not ready"254        print(f"bench:     waiting for server to be ready ...")255        time.sleep(0.5)256 257    print("bench: server started and ready.")258    return server_process259 260 261def start_server_background(args):262    # Start the server263    server_path = '../../../build/bin/llama-server'264    if 'LLAMA_SERVER_BIN_PATH' in os.environ:265        server_path = os.environ['LLAMA_SERVER_BIN_PATH']266    server_args = [267        '--host', args.host,268        '--port', args.port,269    ]270    server_args.extend(['--hf-repo', args.hf_repo])271    server_args.extend(['--hf-file', args.hf_file])272    if args.offline:273        server_args.append('--offline')274    server_args.extend(['--n-gpu-layers', args.n_gpu_layers])275    server_args.extend(['--ctx-size', args.ctx_size])276    server_args.extend(['--parallel', args.parallel])277    server_args.extend(['--batch-size', args.batch_size])278    server_args.extend(['--ubatch-size', args.ubatch_size])279    server_args.extend(['--n-predict', args.max_tokens * 2])280    server_args.append('--cont-batching')281    server_args.append('--metrics')282    server_args.append('--flash-attn')283    args = [str(arg) for arg in [server_path, *server_args]]284    print(f"bench: starting server with: {' '.join(args)}")285    pkwargs = {286        'stdout': subprocess.PIPE,287        'stderr': subprocess.PIPE288    }289    server_process = subprocess.Popen(290        args,291        **pkwargs)  # pyright: ignore[reportArgumentType, reportCallIssue] # ty: ignore[no-matching-overload]292 293    def server_log(in_stream, out_stream):294        for line in iter(in_stream.readline, b''):295            print(line.decode('utf-8'), end='', file=out_stream)296 297    thread_stdout = threading.Thread(target=server_log, args=(server_process.stdout, sys.stdout))298    thread_stdout.start()299    thread_stderr = threading.Thread(target=server_log, args=(server_process.stderr, sys.stderr))300    thread_stderr.start()301 302    return server_process303 304 305def is_server_listening(server_fqdn, server_port):306    with closing(socket.socket(socket.AF_INET, socket.SOCK_STREAM)) as sock:307        result = sock.connect_ex((server_fqdn, server_port))308        _is_server_listening = result == 0309        if _is_server_listening:310            print(f"server is listening on {server_fqdn}:{server_port}...")311        return _is_server_listening312 313 314def is_server_ready(server_fqdn, server_port):315    url = f"http://{server_fqdn}:{server_port}/health"316    response = requests.get(url)317    return response.status_code == 200318 319 320def escape_metric_name(metric_name):321    return re.sub('[^A-Z0-9]', '_', metric_name.upper())322 323 324if __name__ == '__main__':325    main()326