| """Exact OCR anchor์ ์ํ Tray ๊ณ์ฝ์ lattice positive/negative score๋ก ๊ณต๋ ๊ฒ์ฆํ๋ค.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| from datetime import datetime, timezone |
| import json |
| from pathlib import Path |
|
|
| import numpy as np |
|
|
| from math_grid_drawer.research.cross_visual import CrossVisualModel |
| from math_grid_drawer.research.equality_visual import EqualityVisualModel |
| from math_grid_drawer.research.math_tray import fraction_tray_boundary_penalties |
| from math_grid_drawer.research.segmentation_lattice import select_lattice_partition |
| from math_grid_drawer.research.tray_joint import ( |
| TrayJointWeights, |
| adjusted_logits, |
| candidate_signals, |
| component_competition_penalties, |
| local_baseline_boundary_penalties, |
| ) |
| from scripts.crohme_lattice_common import load_cache_for_samples, load_cached_split, writer_fit_validation |
| from scripts.evaluate_crohme_lattice_ocr_fusion import _fit_geometry, _score |
| from scripts.train_crohme_segmentation_lattice_joint_selector import _metrics |
|
|
|
|
| def _prepared_signals( |
| samples: list[dict], cached: list[dict], model, *, strict_equality: bool = False, |
| equality_model: EqualityVisualModel | None = None, cross_model: CrossVisualModel | None = None, |
| cross_gap_ratio: float = 0.40, |
| multistroke_family_boost: float = 6.0, |
| ) -> list[dict]: |
| """ํ์ ๋ณ์: ์๋ณธยทv2 cacheยทgeometry modelยทcross gap. ์๋ ์๋ฆฌ: base fusion๊ณผ ๊ตฌ์กฐ signal์ ํ ๋ฒ ๊ณ์ฐํ๋ค.""" |
|
|
| output = [] |
| for sample, scored in zip(samples, _score(cached, model, ocr_weight=0.5), strict=True): |
| partition = select_lattice_partition(scored["candidates"], scored["logits"], scored["stroke_count"], group_bias=-2.0) |
| tray_signal, symbol_signal, infix_signal = candidate_signals( |
| sample["profiled_strokes"], scored["candidates"], scored["ocr_labels"], |
| scored["features"], partition, ocr_families=scored.get("ocr_families"), |
| strict_equality=strict_equality, equality_model=equality_model, |
| cross_model=cross_model, cross_gap_ratio=cross_gap_ratio, |
| multistroke_family_boost=multistroke_family_boost, |
| ) |
| fraction_penalty = fraction_tray_boundary_penalties(scored["candidates"], sample["profiled_strokes"]) |
| competition_penalty = component_competition_penalties( |
| scored["candidates"], scored["features"], mode="joint", |
| ) |
| raw_local_baseline_penalty = local_baseline_boundary_penalties( |
| sample["profiled_strokes"], scored["candidates"], scored["ocr_labels"], |
| scored["features"], partition, ocr_families=scored.get("ocr_families"), |
| ) |
| local_baseline_penalty = raw_local_baseline_penalty.copy() |
| |
| protected = (tray_signal > 0.0) | (symbol_signal > 0.0) | (infix_signal > 0.0) |
| local_baseline_penalty[protected] = 0.0 |
| output.append({ |
| **scored, "tray_signal": tray_signal, "symbol_signal": symbol_signal, |
| "infix_signal": infix_signal, "fraction_penalty": fraction_penalty, |
| "competition_penalty": competition_penalty, |
| "raw_local_baseline_penalty": raw_local_baseline_penalty, |
| "local_baseline_penalty": local_baseline_penalty, |
| }) |
| return output |
|
|
|
|
| def _weighted( |
| rows: list[dict], *, tray_weight: float, symbol_weight: float, |
| fraction_weight: float, infix_weight: float = 0.0, competition_weight: float = 0.0, |
| local_baseline_weight: float = 0.0, |
| ) -> list[dict]: |
| """ํ์ ๋ณ์: ์ฌ์ ๊ณ์ฐ signalยท๊ฐ์ค์น. ์๋ ์๋ฆฌ: positive ๊ตฌ์กฐ์ ์ธ negative guard๋ฅผ ๊ฒฐํฉํ๋ค.""" |
|
|
| weights = TrayJointWeights( |
| tray=tray_weight, symbol=symbol_weight, fraction=fraction_weight, |
| infix=infix_weight, competition=competition_weight, |
| local_baseline=local_baseline_weight, |
| ) |
| return [{**row, "logits": adjusted_logits( |
| row["logits"], row["tray_signal"], row["symbol_signal"], row["fraction_penalty"], weights, |
| row["infix_signal"], row["competition_penalty"], row["local_baseline_penalty"], |
| )} for row in rows] |
|
|
|
|
| def main() -> None: |
| """ํ์ ๋ณ์: CROHME train/testยทv2 cache. ์๋ ์๋ฆฌ: writer-validation์์ joint ๊ฐ์ค์น๋ฅผ ๊ณ ์ ํ๊ณ official test์ ํ ๋ฒ ์ ์ฉํ๋ค.""" |
|
|
| parser = argparse.ArgumentParser(description="Evaluate CROHME Tray joint selector") |
| parser.add_argument("--train-root", type=Path, required=True) |
| parser.add_argument("--test-root", type=Path, required=True) |
| parser.add_argument("--cache-dir", type=Path, required=True) |
| parser.add_argument("--bundle", type=Path, required=True) |
| parser.add_argument("--output", type=Path, required=True) |
| parser.add_argument("--profile", default="median_height_32") |
| parser.add_argument("--equality-model", type=Path, help="์ ํ์ equality visual JSON head") |
| parser.add_argument("--cross-model", type=Path, help="์ ํ์ cross visual JSON head") |
| parser.add_argument("--cross-gap-ratio", type=float, default=0.40, help="cross head ์ฌ์ bbox gap/์์ ๋์ด ๋น์จ") |
| parser.add_argument("--multistroke-family-boost", type=float, default=6.0, help="OCR family ํธํ ๋คํ ํ๋ณด์ ์ถ๊ฐ symbol signal") |
| parser.add_argument("--fixed-selected", action="store_true", help="๊ธฐ์กด validation ์ ํ๊ฐ 4/4/8/-2๋ฅผ ์ฌ๊ฒ์ฆํ๊ณ sweep์ ์๋ต") |
| parser.add_argument("--sweep-infix", action="store_true", help="๊ธฐ์กด ๊ฐ์ค์น๋ ๊ณ ์ ํ๊ณ x/= ๊ตฌ์กฐ ๊ฐ์ค์น๋ง validation ์ ํ") |
| equality_mode = parser.add_mutually_exclusive_group() |
| equality_mode.add_argument("--strict-equality", dest="strict_equality", action="store_true", help="= ํ๋ณด์ ๋ถ๋ถ์ parser ๊ณ์ฝ ์ ์ฉ") |
| equality_mode.add_argument("--relaxed-equality", dest="strict_equality", action="store_false", help="=๋ฅผ ์์ ์๊ฒฐ์ฑ๊ณผ ๋
๋ฆฝ๋ shape๋ก ์ธ์") |
| parser.set_defaults(strict_equality=False) |
| args = parser.parse_args() |
| equality_model = EqualityVisualModel.load(args.equality_model) if args.equality_model else None |
| cross_model = CrossVisualModel.load(args.cross_model) if args.cross_model else None |
| fit, validation = writer_fit_validation(args.train_root, args.profile) |
| model = _fit_geometry(fit) |
| validation_cache = load_cache_for_samples( |
| validation, args.cache_dir, split="validation", profile=args.profile, bundle=args.bundle, version=2, |
| ) |
| validation_rows = _prepared_signals( |
| validation, validation_cache, model, strict_equality=args.strict_equality, |
| equality_model=equality_model, |
| cross_model=cross_model, cross_gap_ratio=args.cross_gap_ratio, |
| multistroke_family_boost=args.multistroke_family_boost, |
| ) |
| if args.sweep_infix: |
| fixed = TrayJointWeights() |
| trials = [] |
| for infix_weight in (0.0, 0.5, 1.0, 2.0, 4.0, 6.0, 8.0, 12.0, 16.0): |
| trials.append({ |
| "tray_weight": fixed.tray, "symbol_weight": fixed.symbol, |
| "fraction_weight": fixed.fraction, "infix_weight": infix_weight, |
| "group_bias": fixed.group_bias, |
| "metrics": _metrics(_weighted( |
| validation_rows, tray_weight=fixed.tray, symbol_weight=fixed.symbol, |
| fraction_weight=fixed.fraction, infix_weight=infix_weight, |
| ), fixed.group_bias), |
| }) |
| winner = max(trials, key=lambda row: (row["metrics"]["exact_partition"], row["metrics"]["pair_f1"])) |
| elif args.fixed_selected: |
| fixed = TrayJointWeights() |
| winner = { |
| "tray_weight": fixed.tray, "symbol_weight": fixed.symbol, |
| "fraction_weight": fixed.fraction, "infix_weight": fixed.infix, "group_bias": fixed.group_bias, |
| "metrics": _metrics(_weighted( |
| validation_rows, tray_weight=fixed.tray, symbol_weight=fixed.symbol, |
| fraction_weight=fixed.fraction, infix_weight=fixed.infix, |
| ), fixed.group_bias), |
| } |
| else: |
| trials = [] |
| for tray_weight in (0.0, 0.5, 1.0, 2.0, 4.0, 8.0): |
| for symbol_weight in (0.0, 0.5, 1.0, 2.0, 4.0): |
| for fraction_weight in (0.0, 8.0): |
| weighted = _weighted( |
| validation_rows, tray_weight=tray_weight, |
| symbol_weight=symbol_weight, fraction_weight=fraction_weight, |
| ) |
| for bias in (-2.5, -2.0, -1.5, -1.0): |
| trials.append({ |
| "tray_weight": tray_weight, "symbol_weight": symbol_weight, |
| "fraction_weight": fraction_weight, "group_bias": bias, |
| "infix_weight": 0.0, |
| "metrics": _metrics(weighted, bias), |
| }) |
| winner = max(trials, key=lambda row: (row["metrics"]["exact_partition"], row["metrics"]["pair_f1"])) |
| test, test_cache = load_cached_split( |
| args.test_root, args.cache_dir, split="official_test", profile=args.profile, |
| bundle=args.bundle, version=2, |
| ) |
| test_rows = _prepared_signals( |
| test, test_cache, model, strict_equality=args.strict_equality, |
| equality_model=equality_model, |
| cross_model=cross_model, cross_gap_ratio=args.cross_gap_ratio, |
| multistroke_family_boost=args.multistroke_family_boost, |
| ) |
| weighted_test = _weighted( |
| test_rows, tray_weight=float(winner["tray_weight"]), symbol_weight=float(winner["symbol_weight"]), |
| fraction_weight=float(winner["fraction_weight"]), |
| infix_weight=float(winner.get("infix_weight", 0.0)), |
| ) |
| report = { |
| "experiment": "R-CROHME-TRAY-JOINT-SELECTOR-001", |
| "generated_at": datetime.now(timezone.utc).isoformat(), "track": "R_noncommercial_only", |
| "selection_mode": "infix_validation_sweep" if args.sweep_infix else ("fixed_refactor_verification" if args.fixed_selected else "writer_validation_sweep"), |
| "strict_equality": args.strict_equality, |
| "equality_model": str(args.equality_model) if args.equality_model else None, |
| "cross_model": str(args.cross_model) if args.cross_model else None, |
| "selected": winner, "official_test": _metrics(weighted_test, float(winner["group_bias"])), |
| "validation_trials": trials if args.sweep_infix else None, |
| "signal_coverage": { |
| "validation_tray_candidates": int(sum(np.count_nonzero(row["tray_signal"]) for row in validation_rows)), |
| "validation_symbol_candidates": int(sum(np.count_nonzero(row["symbol_signal"]) for row in validation_rows)), |
| "test_tray_candidates": int(sum(np.count_nonzero(row["tray_signal"]) for row in test_rows)), |
| "test_symbol_candidates": int(sum(np.count_nonzero(row["symbol_signal"]) for row in test_rows)), |
| "validation_infix_candidates": int(sum(np.count_nonzero(row["infix_signal"]) for row in validation_rows)), |
| "test_infix_candidates": int(sum(np.count_nonzero(row["infix_signal"]) for row in test_rows)), |
| }, |
| "reference_test": {"exact_partition": 0.4836065574, "pair_f1": 0.86628}, |
| "product_validation": False, |
| "interpretation_limit": "HWRT expanded top-label + CROHME R-track Tray joint selector", |
| } |
| args.output.parent.mkdir(parents=True, exist_ok=True) |
| args.output.write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") |
| print(json.dumps(report, ensure_ascii=False, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|