"""연속 수식 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()