| """통과한 P Formula 3-seed teacher를 하나의 모바일 formula-adapter student로 증류한다.""" |
|
|
| 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 Any, 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.external_corpus import read_jsonl |
| from math_grid_drawer.research.p_formula_dataset06 import ( |
| PFormulaTensorBatch06, |
| materialize_p_formula_split06, |
| p_formula_release_metrics06, |
| p_formula_seed_gate06, |
| ) |
| from math_grid_drawer.research.p_formula_gate06 import audit_p_formula_records06 |
| from math_grid_drawer.research.skeleton_adapter06 import SkeletonTrajectoryAdapter06 |
| from scripts.audit_math_ink_06_case_context import _load_model06 |
| from scripts.train_math_ink_06_formula_adapter import _forward06, _targets06 |
| from scripts.train_math_ink_06_p_formula_adapter import _file_sha25606 |
|
|
|
|
| REQUIRED_SEEDS06 = frozenset({17, 31, 47}) |
|
|
|
|
| def validate_p_formula_distillation_inputs06( |
| summary: dict[str, Any], |
| reports: Sequence[dict[str, Any]], |
| *, |
| data_sha256: str, |
| ) -> list[dict[str, Any]]: |
| """필요 변수: 3-seed summary·teacher report·현재 data hash. 작동 원리: 동일 P corpus와 전 seed 통과를 AND로 검증한다.""" |
|
|
| if summary.get("track") != "P_approved_formula_only": |
| raise ValueError("P Formula distillation에는 P-track summary만 허용합니다.") |
| decision = summary.get("decision") or {} |
| if decision.get("student_distillation_allowed") is not True: |
| raise ValueError("3-seed summary가 student distillation을 허용하지 않았습니다.") |
| if str(summary.get("data_sha256") or "") != data_sha256: |
| raise ValueError("현재 P Formula data SHA-256이 3-seed summary와 다릅니다.") |
| if len(reports) != 3 or {int(report["seed"]) for report in reports} != REQUIRED_SEEDS06: |
| raise ValueError("Teacher report는 seed 17·31·47이 정확히 하나씩 필요합니다.") |
| ordered = sorted(reports, key=lambda report: int(report["seed"])) |
| for report in ordered: |
| if report.get("track") != "P_approved_formula_only": |
| raise ValueError("R-track teacher를 P student에 증류할 수 없습니다.") |
| if str(report.get("data_sha256") or "") != data_sha256: |
| raise ValueError("Teacher report의 P Formula data SHA-256이 다릅니다.") |
| if report.get("seed_gate", {}).get("passed") is not True: |
| raise ValueError(f"seed {report['seed']} teacher가 개별 release gate를 통과하지 않았습니다.") |
| return ordered |
|
|
|
|
| def ensemble_teacher_probability06( |
| logits: Sequence[Tensor], |
| *, |
| temperature: float, |
| ) -> Tensor: |
| """필요 변수: seed별 동일 shape logits·temperature. 작동 원리: logit 평균 대신 확률 평균으로 teacher target을 만든다.""" |
|
|
| if not logits or temperature <= 0.0: |
| raise ValueError("Teacher logit과 양수 temperature가 필요합니다.") |
| shape = logits[0].shape |
| if any(value.shape != shape for value in logits): |
| raise ValueError("Teacher logit shape가 서로 다릅니다.") |
| return torch.stack([ |
| (value / temperature).softmax(dim=1) |
| for value in logits |
| ]).mean(dim=0) |
|
|
|
|
| def _parse_args() -> argparse.Namespace: |
| """필요 변수: P corpus·teacher reports/summary·student main adapter. 작동 원리: fail-closed distillation CLI를 만든다.""" |
|
|
| parser = argparse.ArgumentParser(description="Distill Math Ink 0.6 P formula student") |
| parser.add_argument("--data", type=Path, required=True) |
| parser.add_argument("--teacher-report", type=Path, action="append", required=True) |
| parser.add_argument("--summary", type=Path, required=True) |
| parser.add_argument("--student-adapter", type=Path, required=True) |
| parser.add_argument("--output", type=Path, required=True) |
| parser.add_argument("--seed", type=int, default=17) |
| parser.add_argument("--epochs", type=int, default=20) |
| parser.add_argument("--batch-size", type=int, default=128) |
| parser.add_argument("--learning-rate", type=float, default=4e-4) |
| parser.add_argument("--weight-decay", type=float, default=2e-3) |
| parser.add_argument("--hidden-size", type=int, default=64) |
| parser.add_argument("--temperature", type=float, default=2.0) |
| parser.add_argument("--teacher-exact-weight", type=float, default=0.70) |
| parser.add_argument("--teacher-family-weight", type=float, default=1.00) |
| parser.add_argument("--hard-exact-weight", type=float, default=0.20) |
| parser.add_argument("--hard-family-weight", type=float, default=0.50) |
| parser.add_argument("--patience", type=int, default=5) |
| parser.add_argument("--minimum-independent-sources", type=int, default=2) |
| parser.add_argument("--distillation-regression-maximum-pp", type=float, default=1.0) |
| parser.add_argument("--device", choices=("cuda", "cpu"), default="cuda") |
| return parser.parse_args() |
|
|
|
|
| def _seed06(seed: int) -> None: |
| """필요 변수: student 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 _resolve_project_path06(value: str | Path, *, parent: Path | None = None) -> Path: |
| """필요 변수: checkpoint lineage 경로·선택 report parent. 작동 원리: 절대/상대 경로를 존재하는 실제 파일로 해석한다.""" |
|
|
| path = Path(value) |
| candidates = [path] if path.is_absolute() else [ |
| *((parent / path,) if parent is not None else ()), |
| PROJECT_ROOT / path, |
| ] |
| for candidate in candidates: |
| if candidate.is_file(): |
| return candidate |
| raise FileNotFoundError(f"checkpoint 경로를 찾을 수 없습니다: {value}") |
|
|
|
|
| def _load_teacher06( |
| report: dict[str, Any], |
| report_path: Path, |
| *, |
| device: torch.device, |
| ) -> tuple[Any, torch.nn.Module, SkeletonTrajectoryAdapter06, tuple[str, ...]]: |
| """필요 변수: 통과 teacher report/path. 작동 원리: base→shared online→P formula adapter 순서로 합성한다.""" |
|
|
| formula_checkpoint = _resolve_project_path06( |
| str(report["checkpoint"]), |
| parent=report_path.parent, |
| ) |
| payload = torch.load(formula_checkpoint, map_location="cpu", weights_only=False) |
| if payload.get("track") != "P_approved_formula_only": |
| raise ValueError("P Formula teacher checkpoint track이 올바르지 않습니다.") |
| if payload.get("seed_gate_passed") is not True: |
| raise ValueError("개별 gate를 통과하지 않은 teacher checkpoint입니다.") |
| if str(payload.get("data_sha256") or "") != str(report["data_sha256"]): |
| raise ValueError("Teacher checkpoint/report data SHA-256이 다릅니다.") |
| online_path = _resolve_project_path06(str(payload["online_adapter"])) |
| base_path = _resolve_project_path06(str(payload["base_checkpoint"])) |
| engine, online_adapter = _load_model06(base_path, online_path, device) |
| formula_adapter = SkeletonTrajectoryAdapter06( |
| hidden_size=int(payload["hidden_size"]), |
| ).to(device) |
| formula_adapter.load_state_dict(payload["state_dict"]) |
| engine.model.eval() |
| online_adapter.eval() |
| formula_adapter.eval() |
| return engine, online_adapter, formula_adapter, tuple(str(label) for label in engine.labels) |
|
|
|
|
| def _teacher_targets06( |
| teachers: Sequence[tuple[Any, torch.nn.Module, SkeletonTrajectoryAdapter06]], |
| features: Tensor, |
| *, |
| temperature: float, |
| device: torch.device, |
| batch_size: int, |
| ) -> tuple[Tensor, Tensor]: |
| """필요 변수: 세 teacher·한 split feature. 작동 원리: seed별 exact/family probability를 CPU에서 평균한다.""" |
|
|
| exact_rows, family_rows = [], [] |
| for engine, online_adapter, formula_adapter in teachers: |
| exact, family = _forward06( |
| engine.model, |
| online_adapter, |
| formula_adapter, |
| features, |
| device=device, |
| batch_size=batch_size, |
| ) |
| exact_rows.append(exact) |
| family_rows.append(family) |
| return ( |
| ensemble_teacher_probability06(exact_rows, temperature=temperature), |
| ensemble_teacher_probability06(family_rows, temperature=temperature), |
| ) |
|
|
|
|
| def _balanced_loader06( |
| batch: PFormulaTensorBatch06, |
| tensors: Sequence[Tensor], |
| *, |
| batch_size: int, |
| seed: int, |
| ) -> DataLoader: |
| """필요 변수: P batch·학습 tensor. 작동 원리: source×label 역제곱근 sampler로 distillation batch를 만든다.""" |
|
|
| exact_targets = tensors[0] |
| label_counts = Counter(int(value) for value in exact_targets.tolist()) |
| source_counts = Counter(batch.source_ids) |
| weights = torch.tensor([ |
| 1.0 / ( |
| max(label_counts[int(label)], 1) ** 0.5 |
| * max(source_counts[source], 1) ** 0.5 |
| ) |
| for label, source in zip( |
| exact_targets.tolist(), |
| batch.source_ids, |
| strict=True, |
| ) |
| ]) |
| sampler = WeightedRandomSampler( |
| weights, |
| num_samples=len(weights), |
| replacement=True, |
| generator=torch.Generator().manual_seed(seed), |
| ) |
| return DataLoader( |
| TensorDataset(batch.features, *tensors), |
| batch_size=batch_size, |
| sampler=sampler, |
| ) |
|
|
|
|
| def _release_metrics06( |
| logits: Tensor, |
| targets: Tensor, |
| batch: PFormulaTensorBatch06, |
| labels: Sequence[str], |
| ) -> dict[str, Any]: |
| """필요 변수: student/teacher exact logit·split metadata. 작동 원리: 공통 P release metric을 호출한다.""" |
|
|
| return p_formula_release_metrics06( |
| logits, |
| targets, |
| labels=labels, |
| writer_ids=batch.writer_ids, |
| source_ids=batch.source_ids, |
| timestamp_missing=batch.timestamp_missing, |
| pressure_missing=batch.pressure_missing, |
| ) |
|
|
|
|
| def main() -> None: |
| """필요 변수: 통과한 세 teacher와 동일 P corpus. 작동 원리: 하나의 formula adapter student를 학습하고 test regression을 판정한다.""" |
|
|
| args = _parse_args() |
| device = torch.device(args.device) |
| if device.type == "cuda" and not torch.cuda.is_available(): |
| raise RuntimeError("CUDA distillation을 요청했지만 사용할 수 없습니다.") |
| _seed06(args.seed) |
| data_sha256 = _file_sha25606(args.data) |
| summary = json.loads(args.summary.read_text(encoding="utf-8")) |
| reports = [ |
| json.loads(path.read_text(encoding="utf-8")) |
| for path in args.teacher_report |
| ] |
| ordered_reports = validate_p_formula_distillation_inputs06( |
| summary, |
| reports, |
| data_sha256=data_sha256, |
| ) |
| report_paths = { |
| int(json.loads(path.read_text(encoding="utf-8"))["seed"]): path |
| for path in args.teacher_report |
| } |
| loaded = [ |
| _load_teacher06( |
| report, |
| report_paths[int(report["seed"])], |
| device=device, |
| ) |
| for report in ordered_reports |
| ] |
| label_contracts = {labels for *_modules, labels in loaded} |
| if len(label_contracts) != 1: |
| raise ValueError("세 teacher의 exact vocabulary가 다릅니다.") |
| labels = next(iter(label_contracts)) |
| teachers = [(engine, online, formula) for engine, online, formula, _labels in loaded] |
|
|
| records = list(read_jsonl(args.data)) |
| audit = audit_p_formula_records06( |
| records, |
| minimum_independent_sources=args.minimum_independent_sources, |
| ) |
| if not audit["eligible_for_product_evaluation"]: |
| raise ValueError("현재 P Formula corpus가 product preflight를 통과하지 못했습니다.") |
| split_records = { |
| split: [record for record in records if str(record["split"]) == split] |
| for split in ("training", "validation", "test") |
| } |
| batches = { |
| split: materialize_p_formula_split06(values, allowed_labels=labels) |
| for split, values in split_records.items() |
| } |
| targets = { |
| split: _targets06(batch.truths, labels, loaded[0][0].family_labels) |
| for split, batch in batches.items() |
| } |
| teacher_targets = { |
| split: _teacher_targets06( |
| teachers, |
| batch.features, |
| temperature=args.temperature, |
| device=device, |
| batch_size=args.batch_size, |
| ) |
| for split, batch in batches.items() |
| } |
|
|
| student_adapter_path = _resolve_project_path06(args.student_adapter) |
| student_payload = torch.load( |
| student_adapter_path, |
| map_location="cpu", |
| weights_only=False, |
| ) |
| student_base = _resolve_project_path06(str(student_payload["base_checkpoint"])) |
| student_engine, student_online = _load_model06( |
| student_base, |
| student_adapter_path, |
| device, |
| ) |
| if tuple(str(label) for label in student_engine.labels) != labels: |
| raise ValueError("Student main vocabulary가 teacher와 다릅니다.") |
| for parameter in student_engine.model.parameters(): |
| parameter.requires_grad_(False) |
| for parameter in student_online.parameters(): |
| parameter.requires_grad_(False) |
| student_formula = SkeletonTrajectoryAdapter06( |
| hidden_size=args.hidden_size, |
| ).to(device) |
| train_exact, train_family = targets["training"] |
| train_teacher_exact, train_teacher_family = teacher_targets["training"] |
| loader = _balanced_loader06( |
| batches["training"], |
| ( |
| train_exact, |
| train_family, |
| train_teacher_exact, |
| train_teacher_family, |
| ), |
| batch_size=args.batch_size, |
| seed=args.seed, |
| ) |
| optimizer = torch.optim.AdamW( |
| student_formula.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, |
| ) |
| best_key = (-1.0, -1.0) |
| best_state: dict[str, Tensor] | None = None |
| best_epoch = 0 |
| stale = 0 |
| history = [] |
| for epoch in range(1, args.epochs + 1): |
| student_formula.train() |
| losses = [] |
| for ( |
| features, |
| exact_target, |
| family_target, |
| teacher_exact, |
| teacher_family, |
| ) in loader: |
| features = features.to(device) |
| exact_target = exact_target.to(device) |
| family_target = family_target.to(device) |
| teacher_exact = teacher_exact.to(device) |
| teacher_family = teacher_family.to(device) |
| optimizer.zero_grad(set_to_none=True) |
| with torch.no_grad(): |
| online = student_online(features) |
| exact_logits, family_logits = student_engine.model.classify_trajectory( |
| student_formula(online), |
| ) |
| temperature = args.temperature |
| exact_distill = torch.nn.functional.kl_div( |
| (exact_logits / temperature).log_softmax(dim=1), |
| teacher_exact, |
| reduction="batchmean", |
| ) * temperature ** 2 |
| family_distill = torch.nn.functional.kl_div( |
| (family_logits / temperature).log_softmax(dim=1), |
| teacher_family, |
| reduction="batchmean", |
| ) * temperature ** 2 |
| loss = ( |
| args.teacher_exact_weight * exact_distill |
| + args.teacher_family_weight * family_distill |
| + args.hard_exact_weight |
| * torch.nn.functional.cross_entropy(exact_logits, exact_target) |
| + args.hard_family_weight |
| * torch.nn.functional.cross_entropy(family_logits, family_target) |
| ) |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(student_formula.parameters(), 2.0) |
| optimizer.step() |
| losses.append(float(loss.detach())) |
| scheduler.step() |
| validation_logits = _forward06( |
| student_engine.model, |
| student_online, |
| student_formula, |
| batches["validation"].features, |
| device=device, |
| batch_size=args.batch_size, |
| ) |
| validation_metrics = _release_metrics06( |
| validation_logits[0], |
| targets["validation"][0], |
| batches["validation"], |
| 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["visual_family_top1"]), |
| float(validation_metrics["exact_top1"]), |
| ) |
| if key > best_key: |
| best_key = key |
| best_epoch = epoch |
| best_state = deepcopy({ |
| name: value.detach().cpu() |
| for name, value in student_formula.state_dict().items() |
| }) |
| stale = 0 |
| else: |
| stale += 1 |
| if stale >= args.patience: |
| break |
| if best_state is None: |
| raise RuntimeError("Distilled student checkpoint가 선택되지 않았습니다.") |
| student_formula.load_state_dict(best_state) |
| student_test_logits = _forward06( |
| student_engine.model, |
| student_online, |
| student_formula, |
| batches["test"].features, |
| device=device, |
| batch_size=args.batch_size, |
| )[0] |
| student_test = _release_metrics06( |
| student_test_logits, |
| targets["test"][0], |
| batches["test"], |
| labels, |
| ) |
| teacher_exact_probability = teacher_targets["test"][0] |
| teacher_test = _release_metrics06( |
| teacher_exact_probability.clamp_min(1e-9).log(), |
| targets["test"][0], |
| batches["test"], |
| labels, |
| ) |
| seed_gate = p_formula_seed_gate06(student_test) |
| regression = { |
| metric: ( |
| float(teacher_test[metric]) - float(student_test[metric]) |
| ) * 100.0 |
| for metric in ("exact_top1", "exact_top5", "visual_family_top1") |
| } |
| regression_passed = all( |
| drop <= args.distillation_regression_maximum_pp |
| for drop in regression.values() |
| ) |
| distillation_gate_passed = bool(seed_gate["passed"] and regression_passed) |
| args.output.mkdir(parents=True, exist_ok=True) |
| checkpoint = args.output / "p_formula_student_adapter.pt" |
| torch.save({ |
| "schema": "aiflow-math-ink-06-p-formula-student-v1", |
| "state_dict": best_state, |
| "hidden_size": args.hidden_size, |
| "student_base_checkpoint": str(student_base), |
| "student_online_adapter": str(student_adapter_path), |
| "teacher_seeds": [17, 31, 47], |
| "data_sha256": data_sha256, |
| "selected_epoch": best_epoch, |
| "distillation_gate_passed": distillation_gate_passed, |
| "track": "P_approved_formula_only", |
| "teacher_weights_embedded": False, |
| "litert_exported": False, |
| "product_validation": False, |
| }, checkpoint) |
| report = { |
| "experiment": "P-MATH-INK-06-FORMULA-STUDENT-DISTILL-001", |
| "generated_at": datetime.now(timezone.utc).isoformat(), |
| "student_seed": args.seed, |
| "teacher_seeds": [17, 31, 47], |
| "data": str(args.data), |
| "data_sha256": data_sha256, |
| "preflight": audit, |
| "temperature": args.temperature, |
| "loss_weights": { |
| "teacher_exact": args.teacher_exact_weight, |
| "teacher_family": args.teacher_family_weight, |
| "hard_exact": args.hard_exact_weight, |
| "hard_family": args.hard_family_weight, |
| }, |
| "selected_epoch": best_epoch, |
| "teacher_test": teacher_test, |
| "student_test": student_test, |
| "student_seed_gate": seed_gate, |
| "teacher_to_student_drop_pp": regression, |
| "distillation_regression_maximum_pp": args.distillation_regression_maximum_pp, |
| "distillation_regression_passed": regression_passed, |
| "distillation_gate_passed": distillation_gate_passed, |
| "history": history, |
| "checkpoint": checkpoint.name, |
| "checkpoint_bytes": checkpoint.stat().st_size, |
| "teacher_weights_embedded": False, |
| "track": "P_approved_formula_only", |
| "litert_exported": False, |
| "product_validation": False, |
| "next_gate": "composite torch.export→LiteRT parity→Android 3-tier benchmark", |
| } |
| (args.output / "report.json").write_text( |
| json.dumps(report, ensure_ascii=False, indent=2) + "\n", |
| encoding="utf-8", |
| ) |
| print(json.dumps({ |
| "student_test": student_test, |
| "teacher_test": teacher_test, |
| "distillation_gate_passed": distillation_gate_passed, |
| "checkpoint": str(checkpoint), |
| "product_validation": False, |
| }, ensure_ascii=False, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|