CoolFace
Datasetpublic

echodict/llama.cpp

version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes773downloads
compare-logprobs.py280 linesDownload Raw Back to scripts
1import argparse2import requests3import json4from pathlib import Path5import logging6 7logger = logging.getLogger("compare-logprobs")8logging.basicConfig(level=logging.INFO)9 10 11DESCRIPTION = """12Compare logits between llama.cpp and another inference engine using OpenAI-compatible server endpoints.13 14Unlike compare-logits.py, it allows dumping logits from a hosted API endpoint. Useful when it's not possible to run both models locally.15 16Example usage:17    Step 1: Dump logits from two different servers18        python scripts/compare-logprobs.py dump logits_llama.log http://localhost:8080/v1/completions19        python scripts/compare-logprobs.py dump logits_other.log http://other-engine:8000/v1/completions20 21        (optionally, you can add --api-key <key> if the endpoint requires authentication)22 23    Step 2: Compare the dumped logits24        python scripts/compare-logprobs.py compare logits_llama.log logits_other.log report.md25"""26 27 28def get_remote_corpus(url: str, length: int) -> list[str]:29    response = requests.get(url)30    response.raise_for_status()31    corpus = response.text32    words = [w.strip() for w in corpus.strip().split(" ")]33    words = [w for w in words if "<" not in w] # make sure nothing looks like special tokens34    words = [w for w in words if len(w) > 0]  # filter out empty strings35    while len(words) < length:36        words += words37    return words[:length]38 39 40def dump_logits(41    endpoint: str,42    output_path: Path,43    input_words: list[str],44    pattern: list[tuple[bool, int]],45    api_key=None,46):47    logger.info(f"Dumping logits to {output_path} from endpoint {endpoint}...")48    words = input_words49    curr_text = ""50    n_total = sum(n for get, n in pattern if get)51    n_done = 052    i_cur = 053    i_total = len(words)54    with output_path.open("w") as f:55        for get, n in pattern:56            if not get:57                # skip n words58                for i in range(n):59                    curr_text += words.pop(0) + " "60                    i_cur += 161                continue62            # get n words63            for i in range(n):64                curr_text += words.pop(0) + " "65                payload = {66                    "prompt": curr_text.strip(),67                    "temperature": 0.0,68                    "top_k": 1,69                    "max_tokens": 1,70                    "logprobs": 1,71                    "stream": False,72                }73                response = requests.post(74                    endpoint,75                    json=payload,76                    headers={"Authorization": f"Bearer {api_key}"} if api_key else {},77                )78                response.raise_for_status()79                data = response.json()80                data["__index"] = i_cur  # add index for easier debugging later81                data = json.dumps(data)82                f.write(f"{data}\n")83                n_done += 184                i_cur += 185                logger.info(86                    f"\n\n{data}\n\n[Step: {n_done}/{n_total} | Word: {i_cur}/{i_total}]"87                )88    logger.info(f"Logits dumped to {output_path}")89 90 91def get_token_logprobs(data: dict):92    logprobs = data["choices"][0]["logprobs"]93    if "content" in logprobs:94        # llama.cpp case95        top = logprobs["content"][0]["top_logprobs"][0]96        return top["token"], top["logprob"]97    else:98        # vllm case99        tokens = logprobs["tokens"]100        token_logprobs = logprobs["token_logprobs"]101        return tokens[0], token_logprobs[0]102 103 104def clean_text(text: str) -> str:105    return (106        "'"107        + text.replace("\n", "\\n")108        .replace("\t", "\\t")109        .replace("\r", "\\r")110        .replace("|", "\\|")111        + "'"112    )113 114 115def compare_logits(input1: Path, input2: Path, output_path: Path):116    with input1.open("r") as f1, input2.open("r") as f2, output_path.open("w") as fout:117        lines1 = f1.readlines()118        lines2 = f2.readlines()119 120        tab_header = [121            "idx",122            input1.name,123            "logprob_1",124            input2.name,125            "logprob_2",126            "diff (abs)",127        ]128        tab_entries = []129        tab_max_widths = [len(h) for h in tab_header]130 131        assert len(lines1) == len(132            lines2133        ), "Input files must have the same number of lines."134 135        fout.write("# Logits Comparison Report\n\n")136        for i, (line1, line2) in enumerate(zip(lines1, lines2)):137            if not line1.strip() or not line2.strip():138                continue  # skip empty lines139 140            data1 = json.loads(line1)141            data2 = json.loads(line2)142 143            idx1 = data1.get("__index", -1)144            idx2 = data2.get("__index", -1)145            if idx1 != idx2:146                logger.warning(147                    f"Warning: Mismatched indices at line {i}: {idx1} vs {idx2}"148                )149 150            token1, logprob1 = get_token_logprobs(data1)151            token2, logprob2 = get_token_logprobs(data2)152 153            token1 = clean_text(token1)154            token2 = clean_text(token2)155            abs_diff = abs(logprob1 - logprob2)156 157            tab_entries.append(158                (159                    str(idx1 + 1),160                    token1,161                    f"{logprob1:.4f}",162                    token2,163                    f"{logprob2:.4f}",164                    f"{(abs_diff):.4f}",165                )166            )167 168        for i in range(len(tab_entries)):169            for j in range(len(tab_header)):170                tab_max_widths[j] = max(tab_max_widths[j], len(tab_entries[i][j]))171 172        output = ""173        for j in range(len(tab_header)):174            output += f"| {tab_header[j]:<{tab_max_widths[j]}} "175        output += "|\n"176        for j in range(len(tab_header)):177            output += f"|{'-' * (tab_max_widths[j] + 2)}"178        output += "|\n"179        for entry in tab_entries:180            for j in range(len(tab_header)):181                output += f"| {entry[j]:<{tab_max_widths[j]}} "182            output += "|\n"183 184        logger.info("\n" + output)185        fout.write(output)186        logger.info(f"Report written to {output_path}")187 188 189def parse_pattern(pattern: str) -> list[tuple[bool, int]]:190    parts = pattern.split(",")191    result = []192    for i, part in enumerate(parts):193        n = int(part)194        if i % 2 == 0:195            result.append((True, n))  # get n words196        else:197            result.append((False, n))  # skip n words198    return result199 200 201def parse_args() -> argparse.Namespace:202    parser = argparse.ArgumentParser(203        description=DESCRIPTION, formatter_class=argparse.RawTextHelpFormatter204    )205    subparsers = parser.add_subparsers(206        dest="verb", required=True, help="action to perform"207    )208 209    # dump subcommand210    parser_dump = subparsers.add_parser("dump", help="dump logits from an endpoint")211    parser_dump.add_argument(212        "output", type=Path, help="output path for dumped logits (.log)"213    )214    parser_dump.add_argument(215        "endpoint", type=str, help="OAI-compat /completions endpoint"216    )217    parser_dump.add_argument(218        "--api-key",219        type=str,220        default=None,221        help="API key for authentication (if required)",222    )223    parser_dump.add_argument(224        "--file",225        type=str,226        default="https://raw.githubusercontent.com/ggml-org/llama.cpp/eaba92c3dcc980ebe753348855d4a5d75c069997/tools/server/README.md",227        help="File containing prompt to use instead of the default (can also be an URL)",228    )229    parser_dump.add_argument(230        "--pattern",231        type=str,232        default="10,1000,10,4000,10",233        help="Pattern n_get,n_skip,... where n_get is number of words to get and n_skip is number of words to skip (num of words, NOT num of tokens)",234    )235 236    # compare subcommand237    parser_compare = subparsers.add_parser(238        "compare", help="compare two dumped logits files"239    )240    parser_compare.add_argument("input1", type=Path, help="first input file (.log)")241    parser_compare.add_argument("input2", type=Path, help="second input file (.log)")242    parser_compare.add_argument(243        "output", type=Path, help="output path for comparison report (.md)"244    )245 246    try:247        return parser.parse_args()248    except Exception as e:249        parser.print_help()250        raise e251 252 253def main():254    args = parse_args()255 256    if args.verb == "dump":257        pattern = parse_pattern(args.pattern)258        required_words = sum(n for _, n in pattern)259        if args.file.startswith("http"):260            input_words = get_remote_corpus(args.file, required_words)261            logger.info(f"Fetched {len(input_words)} words from remote {args.file}")262        else:263            with open(args.file, "r") as f:264                input_words = f.read().strip().split(" ")265                input_words = [w for w in input_words if len(w) > 0]  # filter out empty strings266                if len(input_words) < required_words:267                    raise ValueError(268                        f"Input file has only {len(input_words)} words, but pattern requires at least {required_words} words."269                    )270        logger.info(f"Using {len(input_words)} words")271        dump_logits(args.endpoint, args.output, input_words, pattern, args.api_key)272    elif args.verb == "compare":273        compare_logits(args.input1, args.input2, args.output)274    else:275        raise ValueError(f"Unknown verb: {args.verb}")276 277 278if __name__ == "__main__":279    main()280