| """최종 selector에서 구성획보다 약한 merge 후보의 must-not-link weight를 validation 선택한다.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| from datetime import datetime, timezone |
| import json |
| from pathlib import Path |
| import sys |
| from typing import Any |
|
|
| PROJECT_ROOT = Path(__file__).parents[1] |
| SOURCE_ROOT = PROJECT_ROOT / "src" |
| for path in (PROJECT_ROOT, SOURCE_ROOT): |
| if str(path) not in sys.path: |
| sys.path.insert(0, str(path)) |
|
|
| from math_grid_drawer.research.cross_visual import CrossVisualModel |
| from math_grid_drawer.research.equality_visual import EqualityVisualModel |
| 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 |
| from scripts.evaluate_crohme_tray_joint_selector import _prepared_signals, _weighted |
| from scripts.sweep_math_ink_06_multistroke_family_guard import _family_metrics06 |
| from scripts.sweep_math_ink_06_x_grouping_guard import _target_metrics06 |
| from scripts.train_crohme_segmentation_lattice_joint_selector import _metrics |
|
|
|
|
| def _parse_args() -> argparse.Namespace: |
| """필요 변수: CROHME split·cache·full selector head. 작동 원리: test 비개입 competition sweep CLI를 만든다.""" |
|
|
| parser = argparse.ArgumentParser(description="Sweep Math Ink 0.6 component competition guard") |
| parser.add_argument( |
| "--train-root", type=Path, |
| default=PROJECT_ROOT / "research/data/R_noncommercial/ICFHR_package/CROHME2012_data/trainData", |
| ) |
| parser.add_argument( |
| "--test-root", type=Path, |
| default=PROJECT_ROOT / "research/data/R_noncommercial/ICFHR_package/CROHME2012_data/testDataGT", |
| ) |
| parser.add_argument( |
| "--cache-dir", type=Path, |
| default=PROJECT_ROOT / "research/runs/crohme_lattice_ocr_cache_v2_20260722", |
| ) |
| parser.add_argument( |
| "--bundle", type=Path, |
| default=Path(r"research\runs\aiflow_ocr_05_dual_trajectory_3seed_20260720\bundle.manifest.json"), |
| ) |
| parser.add_argument( |
| "--cross-model", type=Path, |
| default=PROJECT_ROOT / "research/runs/crohme_cross_visual_loop3_polyline_20260722/cross_visual.json", |
| ) |
| parser.add_argument( |
| "--equality-model", type=Path, |
| default=PROJECT_ROOT / "research/runs/crohme_equality_visual_loop1_20260722/equality_visual.json", |
| ) |
| parser.add_argument("--profile", default="median_height_32") |
| parser.add_argument("--maximum-x-regression-pp", type=float, default=1.0) |
| parser.add_argument("--maximum-family-regression-pp", type=float, default=2.0) |
| parser.add_argument("--maximum-pair-f1-regression-pp", type=float, default=0.25) |
| parser.add_argument("--output", type=Path, required=True) |
| return parser.parse_args() |
|
|
|
|
| def _evaluate_weight06( |
| samples: list[dict[str, Any]], |
| prepared: list[dict[str, Any]], |
| weight: float, |
| ) -> dict[str, Any]: |
| """필요 변수: 한 split의 cached signal·competition weight. 작동 원리: 전역/x/family 지표를 같은 partition에서 계산한다.""" |
|
|
| weighted = _weighted( |
| prepared, |
| tray_weight=4.0, |
| symbol_weight=4.0, |
| fraction_weight=8.0, |
| infix_weight=8.0, |
| competition_weight=weight, |
| ) |
| return { |
| "component_competition_weight": weight, |
| "global": _metrics(weighted, -2.0), |
| "behavior_targets": _target_metrics06( |
| samples, weighted, group_bias=-2.0, |
| ), |
| "families": _family_metrics06(samples, weighted), |
| } |
|
|
|
|
| def _delta_pp06(candidate: float, reference: float) -> float: |
| """필요 변수: 후보·기준 비율. 작동 원리: 채택 판단용 percentage-point 차이를 반환한다.""" |
|
|
| return (candidate - reference) * 100.0 |
|
|
|
|
| def main() -> None: |
| """필요 변수: writer-validation·official test. 작동 원리: 보호 gate 안 exact winner와 기준만 test에서 비교한다.""" |
|
|
| args = _parse_args() |
| fit, validation = writer_fit_validation(args.train_root, args.profile) |
| geometry_model = _fit_geometry(fit) |
| equality_model = EqualityVisualModel.load(args.equality_model) |
| cross_model = CrossVisualModel.load(args.cross_model) |
| validation_cache = load_cache_for_samples( |
| validation, |
| args.cache_dir, |
| split="validation", |
| profile=args.profile, |
| bundle=args.bundle, |
| version=2, |
| ) |
| validation_prepared = _prepared_signals( |
| validation, |
| validation_cache, |
| geometry_model, |
| equality_model=equality_model, |
| cross_model=cross_model, |
| cross_gap_ratio=0.40, |
| multistroke_family_boost=6.0, |
| ) |
| trials = [ |
| _evaluate_weight06(validation, validation_prepared, weight) |
| for weight in (0.0, 1.0, 2.0, 3.0, 4.0, 6.0, 8.0, 12.0) |
| ] |
| reference = trials[0] |
| minimum_x = ( |
| reference["behavior_targets"]["x"]["grouping_recall"] |
| - args.maximum_x_regression_pp / 100.0 |
| ) |
| minimum_family = ( |
| reference["families"]["grouping_recall"] |
| - args.maximum_family_regression_pp / 100.0 |
| ) |
| minimum_pair_f1 = ( |
| reference["global"]["pair_f1"] |
| - args.maximum_pair_f1_regression_pp / 100.0 |
| ) |
| eligible = [ |
| row for row in trials |
| if ( |
| row["behavior_targets"]["x"]["grouping_recall"] >= minimum_x |
| and row["families"]["grouping_recall"] >= minimum_family |
| and row["global"]["pair_f1"] >= minimum_pair_f1 |
| ) |
| ] |
| winner = max(eligible, key=lambda row: ( |
| row["global"]["exact_partition"], |
| row["global"]["pair_f1"], |
| row["global"]["exact_group_recall"], |
| -row["component_competition_weight"], |
| )) |
| test, test_cache = load_cached_split( |
| args.test_root, |
| args.cache_dir, |
| split="official_test", |
| profile=args.profile, |
| bundle=args.bundle, |
| version=2, |
| ) |
| test_prepared = _prepared_signals( |
| test, |
| test_cache, |
| geometry_model, |
| equality_model=equality_model, |
| cross_model=cross_model, |
| cross_gap_ratio=0.40, |
| multistroke_family_boost=6.0, |
| ) |
| official_reference = _evaluate_weight06(test, test_prepared, 0.0) |
| official_winner = _evaluate_weight06( |
| test, test_prepared, float(winner["component_competition_weight"]), |
| ) |
| deltas = { |
| "exact_partition_pp": _delta_pp06( |
| official_winner["global"]["exact_partition"], |
| official_reference["global"]["exact_partition"], |
| ), |
| "pair_f1_pp": _delta_pp06( |
| official_winner["global"]["pair_f1"], |
| official_reference["global"]["pair_f1"], |
| ), |
| "x_grouping_pp": _delta_pp06( |
| official_winner["behavior_targets"]["x"]["grouping_recall"], |
| official_reference["behavior_targets"]["x"]["grouping_recall"], |
| ), |
| "family_grouping_pp": _delta_pp06( |
| official_winner["families"]["grouping_recall"], |
| official_reference["families"]["grouping_recall"], |
| ), |
| } |
| adopted = bool( |
| float(winner["component_competition_weight"]) > 0.0 |
| and deltas["exact_partition_pp"] > 0.0 |
| and deltas["pair_f1_pp"] >= -args.maximum_pair_f1_regression_pp |
| and deltas["x_grouping_pp"] >= -args.maximum_x_regression_pp |
| and deltas["family_grouping_pp"] >= -args.maximum_family_regression_pp |
| ) |
| report = { |
| "experiment": "R-MATH-INK-06-COMPONENT-COMPETITION-GUARD-001", |
| "generated_at": datetime.now(timezone.utc).isoformat(), |
| "selection_contract": { |
| "split": "CROHME trainData writer-validation only", |
| "cross_gap_ratio": 0.40, |
| "multistroke_family_boost": 6.0, |
| "maximum_x_regression_pp": args.maximum_x_regression_pp, |
| "maximum_family_regression_pp": args.maximum_family_regression_pp, |
| "maximum_pair_f1_regression_pp": args.maximum_pair_f1_regression_pp, |
| }, |
| "reference_validation": reference, |
| "winner_validation": winner, |
| "trials": trials, |
| "official_test_reference": official_reference, |
| "official_test_winner": official_winner, |
| "official_test_deltas": deltas, |
| "decision": { |
| "adopted": adopted, |
| "selected_component_competition_weight": ( |
| float(winner["component_competition_weight"]) if adopted else 0.0 |
| ), |
| "reason": ( |
| "exact partition improves within x/family/pair-F1 guards" |
| if adopted else "official adoption gate failed" |
| ), |
| }, |
| "track": "R_noncommercial_only", |
| "product_validation": False, |
| } |
| 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({key: value for key, value in report.items() if key != "trials"}, ensure_ascii=False, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|