"""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() # 신뢰할 수 있는 완성 다획 기호는 구조 경계와 겹쳐도 기존 positive evidence를 우선한다. 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()