| """연속 수식 group의 baseline 채널 주입 방식을 validation에서 선택하고 official test에 고정 적용한다.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| from datetime import datetime, timezone |
| import json |
| import math |
| from pathlib import Path |
| import sys |
| from typing import Any, Sequence |
|
|
| import torch |
|
|
| 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.trajectory_sequence import visual_label_family |
| from scripts.audit_math_ink_06_case_context import _load_model06 |
| from scripts.crohme_lattice_common import writer_fit_validation |
| from scripts.train_crohme_segmentation_lattice_selector import _samples |
| from scripts.train_math_ink_06_behavior_role import ( |
| TARGET_LABELS06, |
| _canonical_group06, |
| _formula_box06, |
| ) |
| from scripts.train_math_ink_06_skeleton_adapter import _fused_exact_family_logits06 |
|
|
|
|
| CONTEXT_MODES06 = ("formula", "missing", "local_full", "height_only") |
|
|
|
|
| def _parse_args() -> argparse.Namespace: |
| """필요 변수: seed별 adapter·CROHME 공식 split. 작동 원리: validation-only context 계약 선택 CLI를 만든다.""" |
|
|
| parser = argparse.ArgumentParser(description="Audit Math Ink 0.6 formula context contract") |
| parser.add_argument("--adapter", type=Path, action="append", required=True) |
| 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("--profile", default="median_height_32") |
| parser.add_argument("--batch-size", type=int, default=256) |
| parser.add_argument("--device", choices=("cuda", "cpu"), default="cuda") |
| parser.add_argument("--output", type=Path, required=True) |
| return parser.parse_args() |
|
|
|
|
| def apply_formula_context_mode06(features: torch.Tensor, mode: str) -> torch.Tensor: |
| """필요 변수: N×128×19 canonical feature·mode. 작동 원리: glyph shape는 보존하고 baseline 5채널만 교체한다.""" |
|
|
| if mode not in CONTEXT_MODES06: |
| raise ValueError(f"지원하지 않는 formula context mode입니다: {mode}") |
| output = features.clone() |
| if mode == "formula": |
| return output |
| if mode == "missing": |
| output[..., 10:15] = 0.0 |
| elif mode == "local_full": |
| output[..., 10] = 0.0 |
| output[..., 11] = 1.0 |
| output[..., 12] = 1.0 |
| output[..., 13] = 0.5 |
| output[..., 14] = 1.0 |
| else: |
| height = output[..., 12].clone() |
| output[..., 10] = 0.0 |
| output[..., 11] = 0.0 |
| output[..., 12] = height |
| output[..., 13] = 0.0 |
| output[..., 14] = 1.0 |
| return output |
|
|
|
|
| def _materialize_formula_groups06( |
| samples: Sequence[dict[str, Any]], |
| labels: Sequence[str], |
| ) -> tuple[torch.Tensor, list[str], torch.Tensor]: |
| """필요 변수: 공식 formula·model vocabulary. 작동 원리: 정답 group별 sequence와 전체/행동대상 mask를 만든다.""" |
|
|
| allowed = set(str(label) for label in labels) |
| sequences = [] |
| truths: list[str] = [] |
| target_mask = [] |
| for formula in samples: |
| formula_box = _formula_box06(formula["strokes"]) |
| for group, label_value in zip( |
| formula["truth_groups"], formula["truth_labels"], strict=True, |
| ): |
| label = str(label_value) |
| if label not in allowed: |
| continue |
| sequences.append(_canonical_group06(group, formula["strokes"], formula_box)) |
| truths.append(label) |
| target_mask.append(label in TARGET_LABELS06) |
| if not sequences: |
| raise ValueError("지원되는 formula symbol group이 없습니다.") |
| return torch.stack(sequences), truths, torch.tensor(target_mask, dtype=torch.bool) |
|
|
|
|
| def formula_context_metrics06( |
| logits: torch.Tensor, |
| truths: Sequence[str], |
| labels: Sequence[str], |
| target_mask: torch.Tensor, |
| ) -> dict[str, dict[str, float | int]]: |
| """필요 변수: exact logit·정답 label·target mask. 작동 원리: 전체와 행동 형태군 slice를 같은 방식으로 평가한다.""" |
|
|
| if len(logits) != len(truths) or len(target_mask) != len(truths): |
| raise ValueError("formula context 평가 분모가 서로 다릅니다.") |
| predictions = logits.argmax(dim=1).tolist() |
| predicted_labels = [str(labels[index]) for index in predictions] |
|
|
| def summarize(indices: list[int]) -> dict[str, float | int]: |
| """필요 변수: slice index. 작동 원리: exact와 visual-family top-1을 반환한다.""" |
|
|
| exact = sum(predicted_labels[index] == truths[index] for index in indices) |
| family = sum( |
| visual_label_family(predicted_labels[index]) |
| == visual_label_family(str(truths[index])) |
| for index in indices |
| ) |
| return { |
| "samples": len(indices), |
| "exact_top1": exact / max(len(indices), 1), |
| "visual_family_top1": family / max(len(indices), 1), |
| } |
|
|
| all_indices = list(range(len(truths))) |
| target_indices = target_mask.nonzero(as_tuple=False).flatten().tolist() |
| return { |
| "all_supported": summarize(all_indices), |
| "behavior_targets": summarize(target_indices), |
| } |
|
|
|
|
| def formula_context_branch_oracle06( |
| logits_by_mode: Sequence[torch.Tensor], |
| truths: Sequence[str], |
| labels: Sequence[str], |
| target_mask: torch.Tensor, |
| ) -> dict[str, dict[str, float | int]]: |
| """필요 변수: mode별 logit·정답. 작동 원리: 어느 branch든 맞힌 비배포 상한을 전체/행동 slice로 계산한다.""" |
|
|
| if not logits_by_mode or any(len(logits) != len(truths) for logits in logits_by_mode): |
| raise ValueError("context branch oracle 분모가 올바르지 않습니다.") |
| predictions = [logits.argmax(dim=1).tolist() for logits in logits_by_mode] |
|
|
| def summarize(indices: list[int]) -> dict[str, float | int]: |
| """필요 변수: slice index. 작동 원리: exact/visual-family branch oracle을 반환한다.""" |
|
|
| exact = family = 0 |
| for index in indices: |
| truth = str(truths[index]) |
| predicted_labels = [str(labels[row[index]]) for row in predictions] |
| exact += int(any(label == truth for label in predicted_labels)) |
| family += int(any( |
| visual_label_family(label) == visual_label_family(truth) |
| for label in predicted_labels |
| )) |
| return { |
| "samples": len(indices), |
| "exact_oracle": exact / max(len(indices), 1), |
| "visual_family_oracle": family / max(len(indices), 1), |
| } |
|
|
| return { |
| "all_supported": summarize(list(range(len(truths)))), |
| "behavior_targets": summarize( |
| target_mask.nonzero(as_tuple=False).flatten().tolist(), |
| ), |
| } |
|
|
|
|
| def _ensemble_logits06( |
| features: torch.Tensor, |
| adapters: Sequence[Path], |
| *, |
| mode: str, |
| device: torch.device, |
| batch_size: int, |
| ) -> tuple[torch.Tensor, tuple[str, ...]]: |
| """필요 변수: formula feature·seed adapter·context mode. 작동 원리: seed log-probability ensemble을 계산한다.""" |
|
|
| transformed = apply_formula_context_mode06(features, mode) |
| seed_rows = [] |
| canonical_labels: tuple[str, ...] | None = None |
| for adapter_path in adapters: |
| payload = torch.load(adapter_path, map_location="cpu", weights_only=False) |
| base_checkpoint = Path(str(payload["base_checkpoint"])) |
| if not base_checkpoint.is_absolute(): |
| base_checkpoint = PROJECT_ROOT / base_checkpoint |
| engine, adapter = _load_model06(base_checkpoint, adapter_path, device) |
| labels = tuple(str(label) for label in engine.labels) |
| if canonical_labels is not None and labels != canonical_labels: |
| raise ValueError("seed별 vocabulary 순서가 다릅니다.") |
| canonical_labels = labels |
| rows = [] |
| engine.model.eval() |
| adapter.eval() |
| with torch.inference_mode(): |
| for start in range(0, len(transformed), batch_size): |
| batch = transformed[start:start + batch_size].to(device) |
| hypotheses = batch.unsqueeze(1).expand(-1, 4, -1, -1).contiguous() |
| exact, _family = _fused_exact_family_logits06( |
| engine, adapter, hypotheses, |
| ) |
| rows.append(exact.cpu().log_softmax(dim=1)) |
| seed_rows.append(torch.cat(rows)) |
| del engine, adapter |
| if device.type == "cuda": |
| torch.cuda.empty_cache() |
| assert canonical_labels is not None |
| ensemble = torch.logsumexp(torch.stack(seed_rows), dim=0) - math.log(len(seed_rows)) |
| return ensemble, canonical_labels |
|
|
|
|
| def main() -> None: |
| """필요 변수: writer-validation과 official test. 작동 원리: context mode를 validation에서 잠그고 test에는 선택 mode만 적용한다.""" |
|
|
| args = _parse_args() |
| if args.train_root.name.casefold() != "traindata" or args.test_root.name.casefold() != "testdatagt": |
| raise ValueError("CROHME2012 공식 trainData/testDataGT 조합만 허용합니다.") |
| device = torch.device(args.device) |
| if device.type == "cuda" and not torch.cuda.is_available(): |
| raise RuntimeError("CUDA 감사를 요청했지만 사용할 수 없습니다.") |
| first_payload = torch.load(args.adapter[0], map_location="cpu", weights_only=False) |
| first_base = Path(str(first_payload["base_checkpoint"])) |
| if not first_base.is_absolute(): |
| first_base = PROJECT_ROOT / first_base |
| first_engine, _first_adapter = _load_model06(first_base, args.adapter[0], device) |
| labels = first_engine.labels |
| del first_engine, _first_adapter |
| if device.type == "cuda": |
| torch.cuda.empty_cache() |
| _fit, validation_samples = writer_fit_validation(args.train_root, args.profile) |
| test_samples = _samples(args.test_root, args.profile) |
| validation_features, validation_truths, validation_target = _materialize_formula_groups06( |
| validation_samples, labels, |
| ) |
| validation_rows = [] |
| validation_logits_by_mode = [] |
| for mode in CONTEXT_MODES06: |
| logits, current_labels = _ensemble_logits06( |
| validation_features, args.adapter, mode=mode, |
| device=device, batch_size=args.batch_size, |
| ) |
| validation_rows.append({ |
| "mode": mode, |
| **formula_context_metrics06( |
| logits, validation_truths, current_labels, validation_target, |
| ), |
| }) |
| validation_logits_by_mode.append(logits) |
| selected = max( |
| validation_rows, |
| key=lambda row: ( |
| float(row["behavior_targets"]["visual_family_top1"]), |
| float(row["all_supported"]["visual_family_top1"]), |
| float(row["all_supported"]["exact_top1"]), |
| row["mode"] == "formula", |
| ), |
| ) |
| test_features, test_truths, test_target = _materialize_formula_groups06( |
| test_samples, labels, |
| ) |
| test_logits, test_labels = _ensemble_logits06( |
| test_features, args.adapter, mode=str(selected["mode"]), |
| device=device, batch_size=args.batch_size, |
| ) |
| report = { |
| "experiment": "R-MATH-INK-06-FORMULA-CONTEXT-CONTRACT-001", |
| "generated_at": datetime.now(timezone.utc).isoformat(), |
| "split_contract": "CROHME trainData writer-validation mode selection; testDataGT one-shot", |
| "context_modes": CONTEXT_MODES06, |
| "validation_sweep": validation_rows, |
| "selected_mode": str(selected["mode"]), |
| "selected_validation": selected, |
| "validation_branch_oracle": formula_context_branch_oracle06( |
| validation_logits_by_mode, |
| validation_truths, |
| labels, |
| validation_target, |
| ), |
| "official_test": formula_context_metrics06( |
| test_logits, test_truths, test_labels, test_target, |
| ), |
| "track": "R_noncommercial_only", |
| "product_validation": False, |
| "distillation_allowed": 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(report, ensure_ascii=False, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|