| """Formula adapter seed 17ยท31ยท47 ๊ฒฐ๊ณผ๋ฅผ ๋์ผ gate๋ก ์์ฝํ๋ค.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| from datetime import datetime, timezone |
| import json |
| from pathlib import Path |
| from statistics import mean, pstdev |
|
|
|
|
| def _metric06(reports: list[dict], section: str, name: str) -> dict: |
| """ํ์ ๋ณ์: seed reportยทsectionยทmetric. ์๋ ์๋ฆฌ: ๊ฐยทํ๊ท ยทํ์คํธ์ฐจยท์ต์๊ฐ์ ๋ฐํํ๋ค.""" |
|
|
| values = [float(report[section][name]) for report in reports] |
| return { |
| "values": values, |
| "mean": mean(values), |
| "std": pstdev(values), |
| "minimum": min(values), |
| } |
|
|
|
|
| def main() -> None: |
| """ํ์ ๋ณ์: seed report ์ธ ๊ฐ. ์๋ ์๋ฆฌ: validation/test 92% gate์ ์ผ๋ฐํ gap์ ๊ณ ์ ํ๋ค.""" |
|
|
| parser = argparse.ArgumentParser(description="Summarize Math Ink 0.6 formula adapters") |
| parser.add_argument("--report", type=Path, action="append", required=True) |
| parser.add_argument("--output", type=Path, required=True) |
| args = parser.parse_args() |
| reports = [ |
| json.loads(path.read_text(encoding="utf-8")) for path in args.report |
| ] |
| if len(reports) != 3: |
| raise ValueError("formula adapter ์์ฝ์๋ seed report ์ธ ๊ฐ๊ฐ ํ์ํฉ๋๋ค.") |
| validation_visual = _metric06( |
| reports, "selected_validation", "visual_family_top1", |
| ) |
| test_visual = _metric06(reports, "official_test", "visual_family_top1") |
| summary = { |
| "experiment": "R-MATH-INK-06-FORMULA-ADAPTER-3SEED-001", |
| "generated_at": datetime.now(timezone.utc).isoformat(), |
| "reports": [str(path) for path in args.report], |
| "metrics": { |
| "validation_exact_top1": _metric06( |
| reports, "selected_validation", "exact_top1", |
| ), |
| "validation_family_head_top1": _metric06( |
| reports, "selected_validation", "family_head_top1", |
| ), |
| "validation_visual_family_top1": validation_visual, |
| "official_test_exact_top1": _metric06( |
| reports, "official_test", "exact_top1", |
| ), |
| "official_test_family_head_top1": _metric06( |
| reports, "official_test", "family_head_top1", |
| ), |
| "official_test_visual_family_top1": test_visual, |
| "visual_family_generalization_gap_pp": { |
| "values": [ |
| ( |
| float(report["selected_validation"]["visual_family_top1"]) |
| - float(report["official_test"]["visual_family_top1"]) |
| ) * 100.0 |
| for report in reports |
| ], |
| }, |
| }, |
| "decision": { |
| "all_seed_validation_92_passed": validation_visual["minimum"] >= 0.92, |
| "all_seed_official_test_92_passed": test_visual["minimum"] >= 0.92, |
| "architecture_recoverable": validation_visual["minimum"] >= 0.92, |
| "writer_domain_generalization_gate_passed": test_visual["minimum"] >= 0.92, |
| "product_validation": False, |
| "distillation_allowed": False, |
| "next_gate": "P ์ฐ์์ writer/device-disjoint formula adapter ํ์ตยทํ๊ฐ", |
| }, |
| "track": "R_noncommercial_only", |
| "product_validation": False, |
| } |
| gaps = summary["metrics"]["visual_family_generalization_gap_pp"]["values"] |
| summary["metrics"]["visual_family_generalization_gap_pp"].update({ |
| "mean": mean(gaps), |
| "std": pstdev(gaps), |
| "maximum": max(gaps), |
| }) |
| args.output.parent.mkdir(parents=True, exist_ok=True) |
| args.output.write_text( |
| json.dumps(summary, ensure_ascii=False, indent=2) + "\n", |
| encoding="utf-8", |
| ) |
| print(json.dumps(summary, ensure_ascii=False, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|