File size: 12,602 Bytes
4a3cd2f | 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 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 | """연속 수식 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()
|