Felipe97/llama-cpp-compiled
01.2k
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 