echodict/llama.cpp
version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786
0773
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 