aiflow-math-ink-06-intermediate / scripts /sweep_math_ink_06_component_competition_guard.py
cwLeeDev's picture
Add AIFlow Math Ink 0.6 intermediate research snapshot
2948983 verified
Raw
History Blame Contribute Delete
9.08 kB
"""최종 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()