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