| """CROHME truth group에서 연속식 domain residual adapter의 구조적 회복 가능성을 진단한다.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| from collections import Counter |
| from copy import deepcopy |
| from datetime import datetime, timezone |
| import json |
| from pathlib import Path |
| import random |
| import sys |
| from typing import Sequence |
|
|
| import numpy as np |
| import torch |
| from torch import Tensor |
| from torch.utils.data import DataLoader, TensorDataset, WeightedRandomSampler |
|
|
| 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.skeleton_adapter06 import SkeletonTrajectoryAdapter06 |
| from math_grid_drawer.research.trajectory_sequence import shape_family, visual_label_family |
| from scripts.audit_math_ink_06_case_context import _load_model06 |
| from scripts.audit_math_ink_06_formula_context_contract import _materialize_formula_groups06 |
| from scripts.crohme_lattice_common import writer_fit_validation |
| from scripts.train_crohme_segmentation_lattice_selector import _samples |
|
|
|
|
| def _parse_args() -> argparse.Namespace: |
| """필요 변수: 제품 encoder adapter·공식 CROHME split·학습 설정. 작동 원리: R-track shadow adapter CLI를 만든다.""" |
|
|
| parser = argparse.ArgumentParser(description="Train Math Ink 0.6 formula domain adapter") |
| parser.add_argument("--adapter", type=Path, 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("--seed", type=int, default=17) |
| parser.add_argument("--epochs", type=int, default=12) |
| parser.add_argument("--batch-size", type=int, default=256) |
| parser.add_argument("--learning-rate", type=float, default=4e-4) |
| parser.add_argument("--weight-decay", type=float, default=2e-3) |
| parser.add_argument("--exact-loss-weight", type=float, default=0.20) |
| parser.add_argument("--context-dropout", type=float, default=0.20) |
| parser.add_argument("--hidden-size", type=int, default=48) |
| parser.add_argument("--patience", type=int, default=4) |
| parser.add_argument("--skip-official-test", action="store_true") |
| parser.add_argument("--device", choices=("cuda", "cpu"), default="cuda") |
| parser.add_argument("--output", type=Path, required=True) |
| return parser.parse_args() |
|
|
|
|
| def _seed06(seed: int) -> None: |
| """필요 변수: seed. 작동 원리: Python·NumPy·PyTorch 난수를 함께 고정한다.""" |
|
|
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
| if torch.cuda.is_available(): |
| torch.cuda.manual_seed_all(seed) |
|
|
|
|
| def _targets06( |
| truths: Sequence[str], |
| exact_labels: Sequence[str], |
| family_labels: Sequence[str], |
| ) -> tuple[Tensor, Tensor]: |
| """필요 변수: truth label·model ontology. 작동 원리: exact/family target index를 만든다.""" |
|
|
| exact_index = {str(label): index for index, label in enumerate(exact_labels)} |
| family_index = {str(label): index for index, label in enumerate(family_labels)} |
| return ( |
| torch.tensor([exact_index[str(label)] for label in truths], dtype=torch.long), |
| torch.tensor([ |
| family_index[shape_family(str(label))] for label in truths |
| ], dtype=torch.long), |
| ) |
|
|
|
|
| def _metrics06( |
| exact_logits: Tensor, |
| family_logits: Tensor, |
| exact_targets: Tensor, |
| family_targets: Tensor, |
| exact_labels: Sequence[str], |
| ) -> dict[str, float | int]: |
| """필요 변수: exact/family logit·target. 작동 원리: exact·shape-family·visual-family 지표를 같은 분모에서 계산한다.""" |
|
|
| exact_prediction = exact_logits.argmax(dim=1) |
| family_prediction = family_logits.argmax(dim=1) |
| visual_hits = sum( |
| visual_label_family(str(exact_labels[int(prediction)])) |
| == visual_label_family(str(exact_labels[int(truth)])) |
| for prediction, truth in zip( |
| exact_prediction.tolist(), exact_targets.tolist(), strict=True, |
| ) |
| ) |
| return { |
| "samples": len(exact_targets), |
| "exact_top1": float(exact_prediction.eq(exact_targets).float().mean()), |
| "family_head_top1": float(family_prediction.eq(family_targets).float().mean()), |
| "visual_family_top1": visual_hits / max(len(exact_targets), 1), |
| } |
|
|
|
|
| def _forward06( |
| model, |
| online_adapter, |
| formula_adapter, |
| features: Tensor, |
| *, |
| device: torch.device, |
| batch_size: int, |
| ) -> tuple[Tensor, Tensor]: |
| """필요 변수: frozen main·formula adapter·feature. 작동 원리: 전체 split logit을 순서대로 CPU에 반환한다.""" |
|
|
| exact_rows, family_rows = [], [] |
| model.eval() |
| online_adapter.eval() |
| formula_adapter.eval() |
| with torch.inference_mode(): |
| for start in range(0, len(features), batch_size): |
| batch = features[start:start + batch_size].to(device) |
| exact, family = model.classify_trajectory( |
| formula_adapter(online_adapter(batch)), |
| ) |
| exact_rows.append(exact.cpu()) |
| family_rows.append(family.cpu()) |
| return torch.cat(exact_rows), torch.cat(family_rows) |
|
|
|
|
| def _weighted_loader06( |
| dataset: TensorDataset, |
| exact_targets: Tensor, |
| *, |
| batch_size: int, |
| seed: int, |
| ) -> DataLoader: |
| """필요 변수: 학습 tensor·exact target. 작동 원리: label 빈도 역제곱근 sampler로 대형 class 독점을 완화한다.""" |
|
|
| counts = torch.bincount(exact_targets).float() |
| class_weight = counts.clamp_min(1.0).rsqrt() |
| weights = class_weight[exact_targets] |
| sampler = WeightedRandomSampler( |
| weights, num_samples=len(dataset), replacement=True, |
| generator=torch.Generator().manual_seed(seed), |
| ) |
| return DataLoader(dataset, batch_size=batch_size, sampler=sampler) |
|
|
|
|
| def main() -> None: |
| """필요 변수: writer fit/validation·official test. 작동 원리: frozen product encoder 앞 residual만 학습해 구조 상한을 측정한다.""" |
|
|
| args = _parse_args() |
| if args.train_root.name.casefold() != "traindata" or args.test_root.name.casefold() != "testdatagt": |
| raise ValueError("CROHME2012 공식 trainData/testDataGT 조합만 허용합니다.") |
| if not 0.0 <= args.context_dropout <= 1.0: |
| raise ValueError("context dropout은 0~1 범위여야 합니다.") |
| device = torch.device(args.device) |
| if device.type == "cuda" and not torch.cuda.is_available(): |
| raise RuntimeError("CUDA 학습을 요청했지만 사용할 수 없습니다.") |
| _seed06(args.seed) |
| adapter_payload = torch.load(args.adapter, map_location="cpu", weights_only=False) |
| base_checkpoint = Path(str(adapter_payload["base_checkpoint"])) |
| if not base_checkpoint.is_absolute(): |
| base_checkpoint = PROJECT_ROOT / base_checkpoint |
| engine, online_adapter = _load_model06(base_checkpoint, args.adapter, device) |
| for parameter in engine.model.parameters(): |
| parameter.requires_grad_(False) |
| for parameter in online_adapter.parameters(): |
| parameter.requires_grad_(False) |
| fit_samples, validation_samples = writer_fit_validation(args.train_root, args.profile) |
| test_samples = [] if args.skip_official_test else _samples(args.test_root, args.profile) |
| train_x, train_truths, _train_behavior = _materialize_formula_groups06( |
| fit_samples, engine.labels, |
| ) |
| validation_x, validation_truths, _validation_behavior = _materialize_formula_groups06( |
| validation_samples, engine.labels, |
| ) |
| if test_samples: |
| test_x, test_truths, _test_behavior = _materialize_formula_groups06( |
| test_samples, engine.labels, |
| ) |
| else: |
| test_x = torch.empty((0, 128, 19), dtype=train_x.dtype) |
| test_truths = [] |
| train_exact, train_family = _targets06( |
| train_truths, engine.labels, engine.family_labels, |
| ) |
| validation_exact, validation_family = _targets06( |
| validation_truths, engine.labels, engine.family_labels, |
| ) |
| test_exact, test_family = ( |
| _targets06(test_truths, engine.labels, engine.family_labels) |
| if test_truths else ( |
| torch.empty(0, dtype=torch.long), |
| torch.empty(0, dtype=torch.long), |
| ) |
| ) |
| training = TensorDataset(train_x, train_exact, train_family) |
| loader = _weighted_loader06( |
| training, train_exact, batch_size=args.batch_size, seed=args.seed, |
| ) |
| formula_adapter = SkeletonTrajectoryAdapter06(hidden_size=args.hidden_size).to(device) |
| optimizer = torch.optim.AdamW( |
| formula_adapter.parameters(), lr=args.learning_rate, |
| weight_decay=args.weight_decay, |
| ) |
| scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( |
| optimizer, T_max=max(args.epochs, 1), eta_min=args.learning_rate * 0.1, |
| ) |
| baseline_exact, baseline_family = _forward06( |
| engine.model, online_adapter, torch.nn.Identity().to(device), validation_x, |
| device=device, batch_size=args.batch_size, |
| ) |
| baseline_validation = _metrics06( |
| baseline_exact, baseline_family, validation_exact, validation_family, |
| engine.labels, |
| ) |
| best_key = (-1.0, -1.0) |
| best_state = None |
| best_epoch = 0 |
| stale = 0 |
| history = [] |
| for epoch in range(1, args.epochs + 1): |
| formula_adapter.train() |
| losses = [] |
| for features, exact_target, family_target in loader: |
| features = features.to(device) |
| exact_target = exact_target.to(device) |
| family_target = family_target.to(device) |
| if args.context_dropout: |
| drop = torch.rand(len(features), device=device) < args.context_dropout |
| features = features.clone() |
| features[drop, :, 10:15] = 0.0 |
| optimizer.zero_grad(set_to_none=True) |
| with torch.no_grad(): |
| online = online_adapter(features) |
| exact, family = engine.model.classify_trajectory( |
| formula_adapter(online), |
| ) |
| loss = ( |
| torch.nn.functional.cross_entropy(family, family_target) |
| + args.exact_loss_weight |
| * torch.nn.functional.cross_entropy(exact, exact_target) |
| ) |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(formula_adapter.parameters(), 2.0) |
| optimizer.step() |
| losses.append(float(loss.detach())) |
| scheduler.step() |
| validation_logits = _forward06( |
| engine.model, online_adapter, formula_adapter, validation_x, |
| device=device, batch_size=args.batch_size, |
| ) |
| validation_metrics = _metrics06( |
| *validation_logits, validation_exact, validation_family, engine.labels, |
| ) |
| row = { |
| "epoch": epoch, |
| "loss": sum(losses) / max(len(losses), 1), |
| "validation": validation_metrics, |
| } |
| history.append(row) |
| print(json.dumps(row, ensure_ascii=False), flush=True) |
| key = ( |
| float(validation_metrics["family_head_top1"]), |
| float(validation_metrics["visual_family_top1"]), |
| ) |
| if key > best_key: |
| best_key = key |
| best_epoch = epoch |
| best_state = deepcopy({ |
| name: value.detach().cpu() |
| for name, value in formula_adapter.state_dict().items() |
| }) |
| stale = 0 |
| else: |
| stale += 1 |
| if stale >= args.patience: |
| break |
| if best_state is None: |
| raise RuntimeError("formula adapter checkpoint가 선택되지 않았습니다.") |
| formula_adapter.load_state_dict(best_state) |
| selected_validation_logits = _forward06( |
| engine.model, online_adapter, formula_adapter, validation_x, |
| device=device, batch_size=args.batch_size, |
| ) |
| selected_validation = _metrics06( |
| *selected_validation_logits, |
| validation_exact, validation_family, engine.labels, |
| ) |
| if len(test_x): |
| test_logits = _forward06( |
| engine.model, online_adapter, formula_adapter, test_x, |
| device=device, batch_size=args.batch_size, |
| ) |
| official_test = _metrics06( |
| *test_logits, test_exact, test_family, engine.labels, |
| ) |
| else: |
| official_test = None |
| args.output.mkdir(parents=True, exist_ok=True) |
| checkpoint = args.output / "formula_adapter.pt" |
| torch.save({ |
| "schema": "aiflow-math-ink-06-formula-adapter-r-v1", |
| "state_dict": best_state, |
| "hidden_size": args.hidden_size, |
| "base_checkpoint": str(base_checkpoint), |
| "online_adapter": str(args.adapter), |
| "selected_epoch": best_epoch, |
| "context_dropout": args.context_dropout, |
| "exact_loss_weight": args.exact_loss_weight, |
| "track": "R_noncommercial_only", |
| "product_validation": False, |
| "distillation_allowed": False, |
| }, checkpoint) |
| report = { |
| "experiment": "R-MATH-INK-06-FORMULA-ADAPTER-001", |
| "generated_at": datetime.now(timezone.utc).isoformat(), |
| "seed": args.seed, |
| "device": str(device), |
| "cuda_device": ( |
| torch.cuda.get_device_name(device) if device.type == "cuda" else None |
| ), |
| "split_contract": "CROHME trainData writer fit/validation; testDataGT one-shot", |
| "samples": { |
| "fit": len(train_x), |
| "validation": len(validation_x), |
| "official_test": len(test_x), |
| }, |
| "label_support": { |
| "fit": len(Counter(train_truths)), |
| "validation": len(Counter(validation_truths)), |
| "official_test": len(Counter(test_truths)), |
| }, |
| "baseline_validation": baseline_validation, |
| "selected_epoch": best_epoch, |
| "selected_validation": selected_validation, |
| "official_test": official_test, |
| "official_test_skipped": args.skip_official_test, |
| "history": history, |
| "checkpoint": checkpoint.name, |
| "checkpoint_bytes": checkpoint.stat().st_size, |
| "track": "R_noncommercial_only", |
| "product_validation": False, |
| "distillation_allowed": False, |
| } |
| (args.output / "report.json").write_text( |
| json.dumps(report, ensure_ascii=False, indent=2) + "\n", |
| encoding="utf-8", |
| ) |
| print(json.dumps({ |
| "selected_epoch": best_epoch, |
| "baseline_validation": baseline_validation, |
| "selected_validation": selected_validation, |
| "official_test": official_test, |
| "product_validation": False, |
| }, ensure_ascii=False, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|