Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
compare-llama-bench.py381 linesDownload Raw Back to scripts
1#!/usr/bin/env python32 3import logging4import argparse5import heapq6import sys7import os8from glob import glob9import sqlite310 11try:12    import git13    from tabulate import tabulate14except ImportError as e:15    print("the following Python libraries are required: GitPython, tabulate.") # noqa: NP10016    raise e17 18logger = logging.getLogger("compare-llama-bench")19 20# Properties by which to differentiate results per commit:21KEY_PROPERTIES = [22    "cpu_info", "gpu_info", "backends", "n_gpu_layers", "model_filename", "model_type", "n_batch", "n_ubatch",23    "embeddings", "cpu_mask", "cpu_strict", "poll", "n_threads", "type_k", "type_v", "use_mmap", "no_kv_offload",24    "split_mode", "main_gpu", "tensor_split", "flash_attn", "n_prompt", "n_gen"25]26 27# Properties that are boolean and are converted to Yes/No for the table:28BOOL_PROPERTIES = ["embeddings", "cpu_strict", "use_mmap", "no_kv_offload", "flash_attn"]29 30# Header names for the table:31PRETTY_NAMES = {32    "cpu_info": "CPU", "gpu_info": "GPU", "backends": "Backends", "n_gpu_layers": "GPU layers",33    "model_filename": "File", "model_type": "Model", "model_size": "Model size [GiB]",34    "model_n_params": "Num. of par.", "n_batch": "Batch size", "n_ubatch": "Microbatch size",35    "embeddings": "Embeddings", "cpu_mask": "CPU mask", "cpu_strict": "CPU strict", "poll": "Poll",36    "n_threads": "Threads", "type_k": "K type", "type_v": "V type", "split_mode": "Split mode", "main_gpu": "Main GPU",37    "no_kv_offload": "NKVO", "flash_attn": "FlashAttention", "tensor_split": "Tensor split", "use_mmap": "Use mmap",38}39 40DEFAULT_SHOW = ["model_type"]  # Always show these properties by default.41DEFAULT_HIDE = ["model_filename"]  # Always hide these properties by default.42GPU_NAME_STRIP = ["NVIDIA GeForce ", "Tesla ", "AMD Radeon "]  # Strip prefixes for smaller tables.43MODEL_SUFFIX_REPLACE = {" - Small": "_S", " - Medium": "_M", " - Large": "_L"}44 45DESCRIPTION = """Creates tables from llama-bench data written to an SQLite database. Example usage (Linux):46 47$ git checkout master48$ make clean && make llama-bench49$ ./llama-bench -o sql | sqlite3 llama-bench.sqlite50$ git checkout some_branch51$ make clean && make llama-bench52$ ./llama-bench -o sql | sqlite3 llama-bench.sqlite53$ ./scripts/compare-llama-bench.py54 55Performance numbers from multiple runs per commit are averaged WITHOUT being weighted by the --repetitions parameter of llama-bench.56"""57 58parser = argparse.ArgumentParser(59    description=DESCRIPTION, formatter_class=argparse.RawDescriptionHelpFormatter)60help_b = (61    "The baseline commit to compare performance to. "62    "Accepts either a branch name, tag name, or commit hash. "63    "Defaults to latest master commit with data."64)65parser.add_argument("-b", "--baseline", help=help_b)66help_c = (67    "The commit whose performance is to be compared to the baseline. "68    "Accepts either a branch name, tag name, or commit hash. "69    "Defaults to the non-master commit for which llama-bench was run most recently."70)71parser.add_argument("-c", "--compare", help=help_c)72help_i = (73    "Input SQLite file for comparing commits. "74    "Defaults to 'llama-bench.sqlite' in the current working directory. "75    "If no such file is found and there is exactly one .sqlite file in the current directory, "76    "that file is instead used as input."77)78parser.add_argument("-i", "--input", help=help_i)79help_o = (80    "Output format for the table. "81    "Defaults to 'pipe' (GitHub compatible). "82    "Also supports e.g. 'latex' or 'mediawiki'. "83    "See tabulate documentation for full list."84)85parser.add_argument("-o", "--output", help=help_o, default="pipe")86help_s = (87    "Columns to add to the table. "88    "Accepts a comma-separated list of values. "89    f"Legal values: {', '.join(KEY_PROPERTIES[:-2])}. "90    "Defaults to model name (model_type) and CPU and/or GPU name (cpu_info, gpu_info) "91    "plus any column where not all data points are the same. "92    "If the columns are manually specified, then the results for each unique combination of the "93    "specified values are averaged WITHOUT weighing by the --repetitions parameter of llama-bench."94)95parser.add_argument("--check", action="store_true", help="check if all required Python libraries are installed")96parser.add_argument("-s", "--show", help=help_s)97parser.add_argument("--verbose", action="store_true", help="increase output verbosity")98 99known_args, unknown_args = parser.parse_known_args()100 101logging.basicConfig(level=logging.DEBUG if known_args.verbose else logging.INFO)102 103if known_args.check:104    # Check if all required Python libraries are installed. Would have failed earlier if not.105    sys.exit(0)106 107if unknown_args:108    logger.error(f"Received unknown args: {unknown_args}.\n")109    parser.print_help()110    sys.exit(1)111 112input_file = known_args.input113if input_file is None and os.path.exists("./llama-bench.sqlite"):114    input_file = "llama-bench.sqlite"115if input_file is None:116    sqlite_files = glob("*.sqlite")117    if len(sqlite_files) == 1:118        input_file = sqlite_files[0]119 120if input_file is None:121    logger.error("Cannot find a suitable input file, please provide one.\n")122    parser.print_help()123    sys.exit(1)124 125connection = sqlite3.connect(input_file)126cursor = connection.cursor()127builds = cursor.execute("SELECT DISTINCT build_commit FROM test;").fetchall()128 129commit_short_len = len(builds[0][0])130 131try:132    repo = git.Repo(".", search_parent_directories=True)133except git.InvalidGitRepositoryError:134    repo = None135 136 137def find_parent_in_data(commit: git.Commit):138    """Helper function to find the most recent parent measured in number of commits for which there is data."""139    heap: list[tuple[int, git.Commit]] = [(0, commit)]140    seen_hexsha8 = set()141    while heap:142        depth, current_commit = heapq.heappop(heap)143        current_hexsha8 = commit.hexsha[:commit_short_len]144        if (current_hexsha8,) in builds:145            return current_hexsha8146        for parent in commit.parents:147            parent_hexsha8 = parent.hexsha[:commit_short_len]148            if parent_hexsha8 not in seen_hexsha8:149                seen_hexsha8.add(parent_hexsha8)150                heapq.heappush(heap, (depth + 1, parent))151    return None152 153 154def get_all_parent_hexsha8s(commit: git.Commit):155    """Helper function to recursively get hexsha8 values for all parents of a commit."""156    unvisited = [commit]157    visited   = []158 159    while unvisited:160        current_commit = unvisited.pop(0)161        visited.append(current_commit.hexsha[:commit_short_len])162        for parent in current_commit.parents:163            if parent.hexsha[:commit_short_len] not in visited:164                unvisited.append(parent)165 166    return visited167 168 169def get_commit_name(hexsha8):170    """Helper function to find a human-readable name for a commit if possible."""171    if repo is None:172        return hexsha8173    for h in repo.heads:174        if h.commit.hexsha[:commit_short_len] == hexsha8:175            return h.name176    for t in repo.tags:177        if t.commit.hexsha[:commit_short_len] == hexsha8:178            return t.name179    return hexsha8180 181 182def get_commit_hexsha8(name):183    """Helper function to search for a commit given a human-readable name."""184    if repo is None:185        return None186    for h in repo.heads:187        if h.name == name:188            return h.commit.hexsha[:commit_short_len]189    for t in repo.tags:190        if t.name == name:191            return t.commit.hexsha[:commit_short_len]192    for c in repo.iter_commits("--all"):193        if c.hexsha[:commit_short_len] == name[:commit_short_len]:194            return c.hexsha[:commit_short_len]195    return None196 197 198hexsha8_baseline = name_baseline = None199 200# If the user specified a baseline, try to find a commit for it:201if known_args.baseline is not None:202    if (known_args.baseline,) in builds:203        hexsha8_baseline = known_args.baseline204    if hexsha8_baseline is None:205        hexsha8_baseline = get_commit_hexsha8(known_args.baseline)206        name_baseline = known_args.baseline207    if hexsha8_baseline is None:208        logger.error(f"cannot find data for baseline={known_args.baseline}.")209        sys.exit(1)210# Otherwise, search for the most recent parent of master for which there is data:211elif repo is not None:212    hexsha8_baseline = find_parent_in_data(repo.heads.master.commit)213 214    if hexsha8_baseline is None:215        logger.error("No baseline was provided and did not find data for any master branch commits.\n")216        parser.print_help()217        sys.exit(1)218else:219    logger.error("No baseline was provided and the current working directory "220                 "is not part of a git repository from which a baseline could be inferred.\n")221    parser.print_help()222    sys.exit(1)223 224 225name_baseline = get_commit_name(hexsha8_baseline)226 227hexsha8_compare = name_compare = None228 229# If the user has specified a compare value, try to find a corresponding commit:230if known_args.compare is not None:231    if (known_args.compare,) in builds:232        hexsha8_compare = known_args.compare233    if hexsha8_compare is None:234        hexsha8_compare = get_commit_hexsha8(known_args.compare)235        name_compare = known_args.compare236    if hexsha8_compare is None:237        logger.error(f"cannot find data for compare={known_args.compare}.")238        sys.exit(1)239# Otherwise, search for the commit for llama-bench was most recently run240# and that is not a parent of master:241elif repo is not None:242    hexsha8s_master = get_all_parent_hexsha8s(repo.heads.master.commit)243    builds_timestamp = cursor.execute(244        "SELECT build_commit, test_time FROM test ORDER BY test_time;").fetchall()245    for (hexsha8, _) in reversed(builds_timestamp):246        if hexsha8 not in hexsha8s_master:247            hexsha8_compare = hexsha8248            break249 250    if hexsha8_compare is None:251        logger.error("No compare target was provided and did not find data for any non-master commits.\n")252        parser.print_help()253        sys.exit(1)254else:255    logger.error("No compare target was provided and the current working directory "256                 "is not part of a git repository from which a compare target could be inferred.\n")257    parser.print_help()258    sys.exit(1)259 260name_compare = get_commit_name(hexsha8_compare)261 262 263def get_rows(properties):264    """265    Helper function that gets table rows for some list of properties.266    Rows are created by combining those where all provided properties are equal.267    The resulting rows are then grouped by the provided properties and the t/s values are averaged.268    The returned rows are unique in terms of property combinations.269    """270    select_string = ", ".join(271        [f"tb.{p}" for p in properties] + ["tb.n_prompt", "tb.n_gen", "AVG(tb.avg_ts)", "AVG(tc.avg_ts)"])272    equal_string = " AND ".join(273        [f"tb.{p} = tc.{p}" for p in KEY_PROPERTIES] + [274            f"tb.build_commit = '{hexsha8_baseline}'", f"tc.build_commit = '{hexsha8_compare}'"]275    )276    group_order_string = ", ".join([f"tb.{p}" for p in properties] + ["tb.n_gen", "tb.n_prompt"])277    query = (f"SELECT {select_string} FROM test tb JOIN test tc ON {equal_string} "278             f"GROUP BY {group_order_string} ORDER BY {group_order_string};")279    return cursor.execute(query).fetchall()280 281 282# If the user provided columns to group the results by, use them:283if known_args.show is not None:284    show = known_args.show.split(",")285    unknown_cols = []286    for prop in show:287        if prop not in KEY_PROPERTIES[:-2]:  # Last two values are n_prompt, n_gen.288            unknown_cols.append(prop)289    if unknown_cols:290        logger.error(f"Unknown values for --show: {', '.join(unknown_cols)}")291        parser.print_usage()292        sys.exit(1)293    rows_show = get_rows(show)294# Otherwise, select those columns where the values are not all the same:295else:296    rows_full = get_rows(KEY_PROPERTIES)297    properties_different = []298    for i, kp_i in enumerate(KEY_PROPERTIES):299        if kp_i in DEFAULT_SHOW or kp_i == "n_prompt" or kp_i == "n_gen":300            continue301        for row_full in rows_full:302            if row_full[i] != rows_full[0][i]:303                properties_different.append(kp_i)304                break305 306    show = []307    # Show CPU and/or GPU by default even if the hardware for all results is the same:308    if "n_gpu_layers" not in properties_different:309        ngl = int(rows_full[0][KEY_PROPERTIES.index("n_gpu_layers")])310 311        if ngl != 99 and "cpu_info" not in properties_different:312            show.append("cpu_info")313 314    show += properties_different315 316    index_default = 0317    for prop in ["cpu_info", "gpu_info", "n_gpu_layers", "main_gpu"]:318        if prop in show:319            index_default += 1320    show = show[:index_default] + DEFAULT_SHOW + show[index_default:]321    for prop in DEFAULT_HIDE:322        try:323            show.remove(prop)324        except ValueError:325            pass326    rows_show = get_rows(show)327 328table = []329for row in rows_show:330    n_prompt = int(row[-4])331    n_gen    = int(row[-3])332    if n_prompt != 0 and n_gen == 0:333        test_name = f"pp{n_prompt}"334    elif n_prompt == 0 and n_gen != 0:335        test_name = f"tg{n_gen}"336    else:337        test_name = f"pp{n_prompt}+tg{n_gen}"338    #           Regular columns    test name    avg t/s values              Speedup339    #            VVVVVVVVVVVVV     VVVVVVVVV    VVVVVVVVVVVVVV              VVVVVVV340    table.append(list(row[:-4]) + [test_name] + list(row[-2:]) + [float(row[-1]) / float(row[-2])])341 342# Some a-posteriori fixes to make the table contents prettier:343for bool_property in BOOL_PROPERTIES:344    if bool_property in show:345        ip = show.index(bool_property)346        for row_table in table:347            row_table[ip] = "Yes" if int(row_table[ip]) == 1 else "No"348 349if "model_type" in show:350    ip = show.index("model_type")351    for (old, new) in MODEL_SUFFIX_REPLACE.items():352        for row_table in table:353            row_table[ip] = row_table[ip].replace(old, new)354 355if "model_size" in show:356    ip = show.index("model_size")357    for row_table in table:358        row_table[ip] = float(row_table[ip]) / 1024 ** 3359 360if "gpu_info" in show:361    ip = show.index("gpu_info")362    for row_table in table:363        for gns in GPU_NAME_STRIP:364            row_table[ip] = row_table[ip].replace(gns, "")365 366        gpu_names = row_table[ip].split("/")367        num_gpus = len(gpu_names)368        all_names_the_same = len(set(gpu_names)) == 1369        if len(gpu_names) >= 2 and all_names_the_same:370            row_table[ip] = f"{num_gpus}x {gpu_names[0]}"371 372headers  = [PRETTY_NAMES[p] for p in show]373headers += ["Test", f"t/s {name_baseline}", f"t/s {name_compare}", "Speedup"]374 375print(tabulate( # noqa: NP100376    table,377    headers=headers,378    floatfmt=".2f",379    tablefmt=known_args.output380))381