aiflow-math-ink-06-intermediate / scripts /audit_math_ink_06_formula_context_contract.py
cwLeeDev's picture
Add 3-seed formula-domain generalization study
4a3cd2f verified
Raw
History Blame Contribute Delete
12.6 kB
"""연속 수식 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()