"""Evaluate frozen AzooKey candidates, preserving official references and order.""" import argparse import json import math import time from pathlib import Path from vimeml.benchmarks.ajimee import DEVELOPMENT_FORMAT, FORMAT, ROOT, convert_items, json_bytes, sha def read_json(path): return json.loads(path.read_text(encoding="utf-8-sig")) def load_export(directory): manifest = read_json(directory / "manifest.json") if manifest.get("format") not in (FORMAT, DEVELOPMENT_FORMAT) or manifest.get("status") != "complete": raise ValueError("Expected a complete prepared AJIMEE manifest.") for name, digest in manifest["files_sha256"].items(): if sha((directory / name).read_bytes()) != digest: raise ValueError(f"Prepared file changed: {name}.") source = directory / "evaluation_items.json" if sha(source.read_bytes()) != manifest["source_sha256"]: raise ValueError("Source snapshot hash mismatch.") inputs, mapping, stats = convert_items(read_json(source)) if inputs != read_json(directory / "ajimee-input.json") or mapping != read_json(directory / "case-map.json") or stats != manifest["stats"]: raise ValueError("Prepared input/mapping does not match source snapshot.") exported = directory / "ajimee-results" if inputs != read_json(exported / "ajimee-input.json"): raise ValueError("Mac input differs from frozen local input.") raw = read_json(exported / "azookey-candidates.json") items, n_best = raw.get("items"), raw.get("n_best") if type(n_best) is not int or n_best < 1 or not isinstance(items, list) or len(items) != len(inputs): raise ValueError("Invalid n_best or exported sample count.") rows = [] for position, (expected, item, case) in enumerate(zip(inputs, items, mapping)): if any(item.get(key) != expected[value] for key, value in (("query", "query"), ("answers", "answer"), ("left_context", "left_context"))): raise ValueError(f"Row {position}: query/context/references misaligned.") if item.get("right_context") not in (None, ""): raise ValueError(f"Row {position}: unexpected right context.") candidates = item.get("outputs") if not isinstance(candidates, list) or len(candidates) > n_best: raise ValueError(f"Row {position}: invalid candidate count.") texts = [] for candidate in candidates: text, score = candidate.get("text"), candidate.get("score") if not isinstance(text, str) or not text or type(score) not in (int, float) or not math.isfinite(score): raise ValueError(f"Row {position}: invalid candidate text/score.") texts.append(text) if len(set(texts)) != len(texts): raise ValueError(f"Row {position}: duplicate candidate text.") rank = next((i for i, text in enumerate(texts) if text in case["answers"]), -1) if item.get("max_rank") != rank: raise ValueError(f"Row {position}: CLI rank inconsistent with candidates.") rows.append({**case, "candidates": candidates, "orders": {"azookey": texts}, "fallbacks": {}}) # Preserve provenance verbatim; the CLI output does not embed every flag. version_files = ["converter-version.txt", "dictionary-versions.txt", "swift-version.txt"] provenance = {name: (exported / name).read_text(encoding="utf-8-sig") for name in version_files} for line in provenance["dictionary-versions.txt"].splitlines(): if not line.startswith(" "): raise ValueError("Dictionary submodule is uninitialized, conflicted or differs from its pinned commit.") provenance["files_sha256"] = {name: sha((exported / name).read_bytes()) for name in ["ajimee-input.json", "azookey-candidates.json", *version_files]} provenance["n_best"] = n_best provenance["cli_execution_seconds"] = raw.get("execution_time") provenance["flag_provenance"] = "Export requested with typo mode off and no Zenzai; CLI JSON does not embed these flags." return manifest, provenance, rows def edit_distance(reference, hypothesis): previous = list(range(len(hypothesis) + 1)) for i, left in enumerate(reference, 1): current = [i] for j, right in enumerate(hypothesis, 1): current.append(min(current[-1] + 1, previous[j] + 1, previous[j - 1] + (left != right))) previous = current return previous[-1] def min_cer(answers, hypothesis): return min(edit_distance(answer, hypothesis) / len(answer) for answer in answers) def metrics(rows, method): count = len(rows) if not count: return {"cases": 0} top1 = top5 = covered = corrected = regressed = 0 cer = reciprocal = 0.0 for row in rows: order, answers = row["orders"][method], row["answers"] ranks = [i + 1 for i, text in enumerate(order) if text in answers] hit = bool(ranks) and min(ranks) == 1 original = bool(row["orders"]["azookey"]) and row["orders"]["azookey"][0] in answers top1 += hit top5 += bool(ranks) and min(ranks) <= 5 covered += bool(ranks) reciprocal += 1 / min(ranks) if ranks else 0 cer += min_cer(answers, order[0] if order else "") corrected += hit and not original regressed += original and not hit return {"cases": count, "top1_correct": top1, "top1_accuracy": top1 / count, "top5_correct": top5, "top5_accuracy": top5 / count, "candidate_pool_covered": covered, "candidate_pool_recall": covered / count, "covered_top1_accuracy": top1 / covered if covered else None, "mean_min_cer": cer / count, "mrr": reciprocal / count, "corrected_vs_azookey": corrected, "regressed_vs_azookey": regressed, "fallback_cases": sum(method in row["fallbacks"] for row in rows)} def eligibility(lm, context, texts): if not texts: return "empty_candidates" context_ids = lm.processor.encode(context, out_type=int) if lm.processor.decode(context_ids) != context: return "context_roundtrip_mismatch" if len(context_ids) + 1 > lm.model.config.context_length: return "context_length" for text in texts: ids = lm.processor.encode(context + text, out_type=int) if lm.processor.decode(ids) != context + text: return "candidate_roundtrip_mismatch" if not ids or len(ids) > lm.model.config.context_length: return "context_length" return None def rerank(lm, rows): from vimeml.training.evaluate_ime import score_candidates for position, row in enumerate(rows): texts = row["orders"]["azookey"] row["lm_scores"] = {} for field, context, methods in ( ("contextual", row["left_context"], ("lm_context_sum", "lm_context_mean_secondary")), ("context_free", "", ("lm_no_context_sum",)), ): reason = eligibility(lm, context, texts) if reason: for method in methods: row["orders"][method] = list(texts) row["fallbacks"][method] = reason continue scored = score_candidates(lm, context, texts) if any(not math.isfinite(c[key]) for c in scored["candidates"] for key in ("log_probability_sum", "log_probability_mean")): raise ValueError(f"Nonfinite LM score: {row['id']}.") row["lm_scores"][field] = scored for method in methods: key = "log_probability_mean" if method.endswith("secondary") else "log_probability_sum" # Stable sort: ties preserve the actual AzooKey candidate order. row["orders"][method] = [c["text"] for c in sorted(scored["candidates"], key=lambda c: c[key], reverse=True)] if (position + 1) % 25 == 0 or position == len(rows) - 1: print(f"Evaluated {position + 1}/{len(rows)}", flush=True) def summarize(rows): groups = {"all": rows, "with_context": [r for r in rows if r["left_context"]], "without_context": [r for r in rows if not r["left_context"]]} if "lm_context_sum" in rows[0]["orders"]: groups["contextually_scorable"] = [r for r in rows if "lm_context_sum" not in r["fallbacks"]] return {name: {method: metrics(items, method) for method in rows[0]["orders"]} for name, items in groups.items()} def write_summary(path, report, rows): lines = ["# AJIMEE / AzooKey + Tiny LM", "", "完整输入、全部可接受答案、真实候选池;不注入答案,不正规化文字,不加EOS。", "上下文使用官方给定左文,可能跨句;联合分词的公共前缀后logP是候选评分代理,不是精确字符串条件概率。", "主指标为sum,mean仅作预先声明的次要诊断;没有在本评测集调组合系数。并列保持AzooKey原顺序。", "任何候选超出128-token评分窗口时,该策略整条回退原排序;空候选计失败,仍保留完整分母。", "正确答案未进入候选池的样本也保留;训练语料与公开评测文本重叠未知,分句DataLoader的test split没有用于本轮评测。", "", "| 样本组 | 策略 | 样本数 | Top-1 | Top-5 | 候选池覆盖 | MinCER | 纠正 / 改坏 | 回退 |", "| --- | --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: |"] for group, methods in report["metrics"].items(): for method, item in methods.items(): if not item["cases"]: continue lines.append(f"| {group} | {method} | {item['cases']} | {item['top1_correct']} ({item['top1_accuracy']:.1%}) | " f"{item['top5_correct']} ({item['top5_accuracy']:.1%}) | {item['candidate_pool_covered']} ({item['candidate_pool_recall']:.1%}) | " f"{item['mean_min_cer']:.4f} | {item['corrected_vs_azookey']} / {item['regressed_vs_azookey']} | {item['fallback_cases']} |") lines += ["", "候选池覆盖是这次冻结候选下精确命中的上限,不等于理论上能达到的实际模型性能。开发机耗时不是iOS延迟。", "候选原分数与全部LM token分数见scores.jsonl;changed-cases.json记录纠正/改坏样本;metrics.json记录输入、候选、版本和模型hash。", ""] path.write_text("\n".join(lines), encoding="utf-8") def main(argv=None): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--benchmark", type=Path, default=ROOT / "artifacts/benchmarks/ajimee-jwtd-v2-v1") parser.add_argument("--checkpoint", type=Path, default=ROOT / "artifacts/models/tiny-ja-v1/best.pt") parser.add_argument("--tokenizer", type=Path, default=ROOT / "artifacts/tokenizers/ja-unigram-16k-v1") parser.add_argument("--output", type=Path, default=ROOT / "outputs/ime-eval/tiny-ja-v1-ajimee") parser.add_argument("--device", choices=("cpu", "cuda"), default="cpu") parser.add_argument("--threads", type=int, default=4) parser.add_argument("--baseline-only", action="store_true") args = parser.parse_args(argv) if args.threads < 1: parser.error("threads must be positive.") if args.output.exists() and any(args.output.iterdir()): parser.error("Output is not empty; choose a new --output directory to preserve previous results.") started = time.perf_counter() manifest, provenance, rows = load_export(args.benchmark) model = None if not args.baseline_only: import torch from vimeml.training.infer import JapaneseLM torch.set_num_threads(args.threads) lm = JapaneseLM(args.checkpoint, args.tokenizer, args.device) model = lm.metadata rerank(lm, rows) report = {"format": "ajimee_ime_metrics_v1", "status": "complete", "model": model, "benchmark_manifest": manifest, "export_provenance": provenance, "metrics": summarize(rows), "policy": "Original queries/references/candidates; supplied context; suffix logP sum primary; mean secondary; stable ties; whole-case fallback; no EOS or truncation.", "training_overlap": "Unknown; no corpus overlap audit performed.", "test_split_used": False, "empty_candidate_cases": sum(not r["candidates"] for r in rows), "elapsed_seconds": time.perf_counter() - started} changed = [] for row in rows: original = row["orders"]["azookey"] before = bool(original) and original[0] in row["answers"] for method, order in row["orders"].items(): after = bool(order) and order[0] in row["answers"] if before != after: changed.append({"id": row["id"], "method": method, "change": "corrected" if after else "regressed", "query": row["query"], "context": row["left_context"], "answers": row["answers"], "before": original[0], "after": order[0]}) args.output.mkdir(parents=True, exist_ok=True) (args.output / "scores.jsonl").write_text("".join(json.dumps(row, ensure_ascii=False) + "\n" for row in rows), encoding="utf-8") (args.output / "changed-cases.json").write_bytes(json_bytes(changed)) report["files_sha256"] = {name: sha((args.output / name).read_bytes()) for name in ("scores.jsonl", "changed-cases.json")} (args.output / "metrics.json").write_bytes(json_bytes(report)) write_summary(args.output / "results.md", report, rows) print(json.dumps(report["metrics"]["all"], ensure_ascii=False, indent=2)) print(f"Report: {args.output.resolve() / 'results.md'}") if __name__ == "__main__": main()