File size: 4,822 Bytes
f770448
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0eef691
f770448
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0eef691
f770448
 
0eef691
 
 
f770448
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
"""기존과 federated 0.6을 전체 P-source holdout에서 동일 분모로 비교한다."""

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

PROJECT_ROOT = Path(__file__).parents[1]
SOURCE_ROOT = PROJECT_ROOT / "src"
if str(SOURCE_ROOT) not in sys.path:
    sys.path.insert(0, str(SOURCE_ROOT))
if str(PROJECT_ROOT / "scripts") not in sys.path:
    sys.path.insert(0, str(PROJECT_ROOT / "scripts"))

from math_grid_drawer.research.ink06_federation import load_product_federation06, resolve_training_device06
from math_grid_drawer.research.math_ink_06 import MathInk06Engine
from train_math_ink_06_federated_online import _evaluate_source


def main() -> None:
    """필요 변수: base/candidate·세 source. 작동 원리: UCI 공식 test와 HWRT train-writer holdout 전체를 비교한다."""

    parser = argparse.ArgumentParser(description="Evaluate federated Math Ink 0.6")
    parser.add_argument("--baseline", type=Path, required=True)
    parser.add_argument("--candidate", type=Path, required=True)
    parser.add_argument("--registry", type=Path, default=PROJECT_ROOT / "research/dataset_registry.json")
    parser.add_argument("--source-registry", type=Path, default=PROJECT_ROOT / "research/math_ink_06_source_registry.json")
    parser.add_argument("--commercial", type=Path, default=PROJECT_ROOT / "research/data/external_trajectory_v1/commercial_ccby4.jsonl.gz")
    parser.add_argument("--hwrt", type=Path, default=PROJECT_ROOT / "research/data/open_pretrain/hwrt_expanded_v2/hwrt_expanded.jsonl.gz")
    parser.add_argument("--approval", type=Path, default=PROJECT_ROOT / "research/approvals/HWRT-ODBL-USE-APPROVAL-v1.json")
    parser.add_argument("--output", type=Path, required=True)
    parser.add_argument("--batch-size", type=int, default=64)
    parser.add_argument("--source", action="append", default=[], help="지정한 source만 평가한다. 반복 사용 가능")
    parser.add_argument("--device", default="auto", help="auto|cpu|cuda[:index]")
    args = parser.parse_args()

    device = resolve_training_device06(args.device)
    baseline = MathInk06Engine(args.baseline, device=device)
    candidate = MathInk06Engine(args.candidate, device=device)
    if baseline.labels != candidate.labels:
        raise ValueError("baseline과 candidate vocabulary가 다릅니다.")
    exact_to_index = {label: index for index, label in enumerate(candidate.labels)}
    family_to_index = {label: index for index, label in enumerate(candidate.family_labels)}
    sources = load_product_federation06(
        registry_path=args.registry, commercial_path=args.commercial, hwrt_path=args.hwrt,
        approval_path=args.approval, allowed_labels=candidate.labels, source_registry_path=args.source_registry,
    )
    groups = {}
    contracts = {}
    for source in sources:
        if args.source and source.source_id not in args.source:
            continue
        if source.source_id == "hwrt":
            groups[source.source_id] = [row for row in source.records if row.get("split") == "test"]
            contracts[source.source_id] = "HWRT approved train의 AIFlow writer-disjoint test; official test excluded"
        elif source.source_id == "uci-uji-pen-v1":
            groups[source.source_id] = [row for row in source.records if row.get("split") == "test"]
            contracts[source.source_id] = "AIFlow project writer-disjoint w10-w11 test; official UCI task is 11-fold LOOW"
        else:
            groups[source.source_id] = [row for row in source.records if row.get("split") == "test"]
            contracts[source.source_id] = "official UCI writer-independent test"
    report = {
        "baseline": str(args.baseline), "candidate": str(args.candidate), "sources": {},
        "hwrt_official_test_used": False, "product_validation": False,
    }
    for source_id, records in groups.items():
        base_metrics = _evaluate_source(
            baseline, records, exact_to_index, family_to_index, args.batch_size,
        )
        candidate_metrics = _evaluate_source(
            candidate, records, exact_to_index, family_to_index, args.batch_size,
        )
        report["sources"][source_id] = {
            "contract": contracts[source_id], "baseline": base_metrics, "candidate": candidate_metrics,
            "delta": {
                metric: candidate_metrics[metric] - base_metrics[metric]
                for metric in ("online_top1", "online_top5", "raster_top1", "raster_top5")
            },
        }
    args.output.mkdir(parents=True, exist_ok=True)
    (args.output / "federated_evaluation.json").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()