| """Validation에서 exact/family logit 결합 가중치를 고정하고 paired-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 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 scripts.train_math_ink_06_p_boundary_auxiliary import _load_encoder06 |
| from scripts.train_math_ink_06_skeleton_adapter import _resolve_device06 |
| from math_grid_drawer.research.trajectory_sequence import shape_family |
|
|
|
|
| def _parse_args() -> argparse.Namespace: |
| """필요 변수: validation/test cache와 세 seed. 작동 원리: test 선택을 금지한 calibration CLI를 만든다.""" |
|
|
| parser = argparse.ArgumentParser(description="Calibrate Math Ink 0.6 online family fusion") |
| parser.add_argument("--validation-cache", type=Path, required=True) |
| parser.add_argument("--test-cache", type=Path, required=True) |
| parser.add_argument("--base-checkpoint", type=Path, action="append", required=True) |
| parser.add_argument("--adapter-checkpoint", type=Path, action="append", required=True) |
| parser.add_argument("--weights", default="0,0.025,0.05,0.075,0.1,0.125,0.15,0.2,0.25,0.3") |
| parser.add_argument("--batch-size", type=int, default=256) |
| parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto") |
| parser.add_argument("--output", type=Path, required=True) |
| args = parser.parse_args() |
| if len(args.base_checkpoint) != len(args.adapter_checkpoint): |
| raise ValueError("base와 adapter checkpoint 개수는 같아야 합니다.") |
| if len(args.base_checkpoint) < 2: |
| raise ValueError("fusion calibration에는 seed 두 개 이상이 필요합니다.") |
| return args |
|
|
|
|
| def _macro_f106(targets: torch.Tensor, predictions: torch.Tensor) -> float: |
| """필요 변수: 정답·예측 index. 작동 원리: test에 없는 class를 분모에서 제외한 macro-F1을 계산한다.""" |
|
|
| values = [] |
| for label in targets.unique().tolist(): |
| truth = targets.eq(label) |
| predicted = predictions.eq(label) |
| true_positive = int((truth & predicted).sum()) |
| denominator = 2 * true_positive + int((truth & ~predicted).sum()) + int((~truth & predicted).sum()) |
| values.append(2 * true_positive / denominator if denominator else 0.0) |
| return sum(values) / max(len(values), 1) |
|
|
|
|
| def _metrics06(logits: torch.Tensor, targets: torch.Tensor) -> dict[str, float | int]: |
| """필요 변수: fused logit·정답. 작동 원리: 동일 분모의 top-1/top-5/macro-F1을 반환한다.""" |
|
|
| prediction = logits.argmax(dim=-1) |
| top5 = logits.topk(min(5, logits.shape[-1]), dim=-1).indices |
| return { |
| "samples": len(targets), |
| "top1": float(prediction.eq(targets).float().mean()), |
| "top5": float(top5.eq(targets[:, None]).any(dim=-1).float().mean()), |
| "macro_f1": _macro_f106(targets, prediction), |
| } |
|
|
|
|
| def fusion_sweep06( |
| exact_by_seed: Sequence[torch.Tensor], |
| family_by_seed: Sequence[torch.Tensor], |
| targets: torch.Tensor, |
| exact_family_index: torch.Tensor, |
| weights: Sequence[float], |
| ) -> list[dict[str, float | int]]: |
| """필요 변수: seed별 exact/family logit·가중치. 작동 원리: 확률공간 seed ensemble을 weight별 평가한다.""" |
|
|
| if len(exact_by_seed) != len(family_by_seed) or not exact_by_seed: |
| raise ValueError("exact/family seed 출력 개수가 올바르지 않습니다.") |
| rows = [] |
| for weight in weights: |
| seed_joint = [] |
| for exact, family in zip(exact_by_seed, family_by_seed, strict=True): |
| joint = exact.log_softmax(dim=-1) |
| if weight: |
| joint = joint + float(weight) * family.log_softmax(dim=-1)[:, exact_family_index] |
| seed_joint.append(joint) |
| ensemble = torch.logsumexp(torch.stack(seed_joint), dim=0) - math.log(len(seed_joint)) |
| rows.append({"family_fusion_weight": float(weight), **_metrics06(ensemble, targets)}) |
| return rows |
|
|
|
|
| def _infer_split06( |
| cache_path: Path, |
| base_paths: Sequence[Path], |
| adapter_paths: Sequence[Path], |
| *, |
| device: torch.device, |
| batch_size: int, |
| ) -> tuple[list[torch.Tensor], list[torch.Tensor], torch.Tensor, torch.Tensor]: |
| """필요 변수: split cache·composite seed. 작동 원리: fusion 전 exact/family logit과 사상을 수집한다.""" |
|
|
| cache = torch.load(cache_path, map_location="cpu", weights_only=True, mmap=True) |
| features = cache["features"][:, 0] |
| targets = cache["targets"].long().clone() |
| exact_rows, family_rows = [], [] |
| family_index: torch.Tensor | None = None |
| for base_path, adapter_path in zip(base_paths, adapter_paths, strict=True): |
| model, adapter, base, _adapter_payload = _load_encoder06(base_path, adapter_path, device) |
| family_to_index = { |
| str(label): index for index, label in enumerate(base["family_labels"]) |
| } |
| current_family_index = torch.tensor([ |
| family_to_index[shape_family(str(label))] |
| for label in base["exact_labels"] |
| ], dtype=torch.long) |
| if family_index is not None and not torch.equal(family_index, current_family_index): |
| raise ValueError("seed별 exact→family 사상이 다릅니다.") |
| family_index = current_family_index |
| exact_batches, family_batches = [], [] |
| model.eval() |
| adapter.eval() |
| with torch.inference_mode(): |
| for start in range(0, len(features), batch_size): |
| exact, family = model.forward_online( |
| adapter(features[start:start + batch_size].to(device)), |
| ) |
| exact_batches.append(exact.cpu()) |
| family_batches.append(family.cpu()) |
| exact_rows.append(torch.cat(exact_batches)) |
| family_rows.append(torch.cat(family_batches)) |
| del model, adapter |
| if device.type == "cuda": |
| torch.cuda.empty_cache() |
| assert family_index is not None |
| return exact_rows, family_rows, targets, family_index |
|
|
|
|
| def main() -> None: |
| """필요 변수: CLI 설정. 작동 원리: validation winner만 test에 적용하고 결과를 UTF-8 JSON으로 고정한다.""" |
|
|
| args = _parse_args() |
| weights = tuple(float(value.strip()) for value in args.weights.split(",") if value.strip()) |
| if not weights or any(weight < 0.0 or weight > 1.0 for weight in weights): |
| raise ValueError("family fusion weight는 0~1 범위여야 합니다.") |
| device = _resolve_device06(args.device) |
| validation = _infer_split06( |
| args.validation_cache, args.base_checkpoint, args.adapter_checkpoint, |
| device=device, batch_size=args.batch_size, |
| ) |
| validation_sweep = fusion_sweep06(*validation, weights) |
| selected = max( |
| validation_sweep, |
| key=lambda row: (float(row["top1"]), float(row["macro_f1"]), -float(row["family_fusion_weight"])), |
| ) |
| test = _infer_split06( |
| args.test_cache, args.base_checkpoint, args.adapter_checkpoint, |
| device=device, batch_size=args.batch_size, |
| ) |
| test_result = fusion_sweep06( |
| *test, (float(selected["family_fusion_weight"]),), |
| )[0] |
| zero_test = fusion_sweep06(*test, (0.0,))[0] |
| report = { |
| "experiment": "MATH-INK-06-ONLINE-FAMILY-FUSION-CALIBRATION-001", |
| "generated_at": datetime.now(timezone.utc).isoformat(), |
| "selection_contract": "validation only; paired-test evaluated once after weight lock", |
| "device": str(device), |
| "validation_sweep": validation_sweep, |
| "selected_validation": selected, |
| "paired_test_zero_weight": zero_test, |
| "paired_test_selected_weight": test_result, |
| "paired_test_gain_pp": { |
| "top1": (float(test_result["top1"]) - float(zero_test["top1"])) * 100.0, |
| "top5": (float(test_result["top5"]) - float(zero_test["top5"])) * 100.0, |
| "macro_f1": (float(test_result["macro_f1"]) - float(zero_test["macro_f1"])) * 100.0, |
| }, |
| "product_validation": 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() |
|
|