Download scripts/export_math_ink_06_p_formula_student.py from cwLeeDev/aiflow-math-ink-06-intermediate: direct link, hf CLI and curl.
- Browser
- Download file 8.85 kB
-
https://huggingface.co/cwLeeDev/aiflow-math-ink-06-intermediate/resolve/c228b1d168b7f7d1dc45d6c487951b9b42687833/scripts/export_math_ink_06_p_formula_student.py
- Command line
-
hf download hf://cwLeeDev/aiflow-math-ink-06-intermediate@c228b1d168b7f7d1dc45d6c487951b9b42687833/scripts/export_math_ink_06_p_formula_student.py
-
curl -L -o export_math_ink_06_p_formula_student.py https://huggingface.co/cwLeeDev/aiflow-math-ink-06-intermediate/resolve/c228b1d168b7f7d1dc45d6c487951b9b42687833/scripts/export_math_ink_06_p_formula_student.py
8.85 kB
| """통과한 P Formula student를 strict torch.export와 선택적 LiteRT로 고정한다.""" | |
| from __future__ import annotations | |
| import argparse | |
| import importlib.util | |
| import json | |
| from pathlib import Path | |
| import sys | |
| from typing import Any | |
| 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.external_corpus import read_jsonl | |
| from math_grid_drawer.research.ink06_export import ( | |
| PFormulaStudentExportWrapper06, | |
| exported_equivalence06, | |
| ) | |
| from math_grid_drawer.research.math_ink_06 import MathInk06Engine | |
| from math_grid_drawer.research.p_formula_dataset06 import materialize_p_formula_split06 | |
| from math_grid_drawer.research.p_formula_gate06 import audit_p_formula_records06 | |
| from math_grid_drawer.research.skeleton_adapter06 import ( | |
| DualModalityTrajectoryAdapter06, | |
| SkeletonTrajectoryAdapter06, | |
| ) | |
| from scripts.export_math_ink_06_litert import ( | |
| _convert_litert, | |
| _save_exported_program06, | |
| _vocabulary_sha25606, | |
| ) | |
| from scripts.train_math_ink_06_p_formula_adapter import _file_sha25606 | |
| MAXIMUM_MODEL_BYTES06 = 25 * 1024 * 1024 | |
| def validate_p_formula_student_export06( | |
| payload: dict[str, Any], | |
| *, | |
| data_sha256: str, | |
| ) -> None: | |
| """필요 변수: student checkpoint metadata·현재 P corpus hash. 작동 원리: 실패한/오염된 student의 export를 차단한다.""" | |
| if payload.get("schema") != "aiflow-math-ink-06-p-formula-student-v1": | |
| raise ValueError("지원하지 않는 P Formula student checkpoint입니다.") | |
| if payload.get("track") != "P_approved_formula_only": | |
| raise ValueError("P 승인 track이 아닌 student는 export할 수 없습니다.") | |
| if payload.get("distillation_gate_passed") is not True: | |
| raise ValueError("정식 distillation gate를 통과하지 않은 student입니다.") | |
| if payload.get("teacher_weights_embedded") is not False: | |
| raise ValueError("Teacher weight가 포함되었거나 포함 여부가 불명확합니다.") | |
| if str(payload.get("data_sha256") or "") != data_sha256: | |
| raise ValueError("Student와 현재 P Formula corpus의 SHA-256이 다릅니다.") | |
| if set(int(seed) for seed in payload.get("teacher_seeds", [])) != {17, 31, 47}: | |
| raise ValueError("Student lineage에는 teacher seed 17·31·47이 모두 필요합니다.") | |
| def _resolve_checkpoint_path06(value: str | Path, *, parent: Path) -> Path: | |
| """필요 변수: checkpoint lineage 값·student parent. 작동 원리: 상대 경로를 제한된 후보에서만 실제 파일로 해석한다.""" | |
| path = Path(value) | |
| candidates = [path] if path.is_absolute() else [parent / path, PROJECT_ROOT / path] | |
| for candidate in candidates: | |
| if candidate.is_file(): | |
| return candidate | |
| raise FileNotFoundError(f"checkpoint lineage 파일을 찾을 수 없습니다: {value}") | |
| def _online_branch06(adapter: torch.nn.Module) -> torch.nn.Module: | |
| """필요 변수: single/dual online adapter. 작동 원리: formula student가 학습에 사용한 online branch만 고정한다.""" | |
| return adapter.online if isinstance(adapter, DualModalityTrajectoryAdapter06) else adapter | |
| def _representative_inputs06( | |
| data: Path, | |
| *, | |
| labels: tuple[str, ...], | |
| maximum_samples: int, | |
| ) -> tuple[list[tuple[torch.Tensor, ...]], dict[str, Any]]: | |
| """필요 변수: 동일 P corpus·378 vocabulary·표본 상한. 작동 원리: test split의 실제 formula-relative tensor를 대표 입력으로 만든다.""" | |
| records = list(read_jsonl(data)) | |
| audit = audit_p_formula_records06(records) | |
| if not audit["eligible_for_product_evaluation"]: | |
| raise ValueError("P Formula corpus가 product preflight를 통과하지 못했습니다.") | |
| test_records = [record for record in records if str(record["split"]) == "test"] | |
| batch = materialize_p_formula_split06(test_records, allowed_labels=labels) | |
| if not len(batch.features): | |
| raise ValueError("Student export에는 test representative가 필요합니다.") | |
| limit = min(len(batch.features), maximum_samples) | |
| return ( | |
| [(batch.features[index:index + 1],) for index in range(limit)], | |
| audit, | |
| ) | |
| def _parse_args() -> argparse.Namespace: | |
| """필요 변수: 통과 student·동일 P corpus·출력. 작동 원리: export와 선택적 LiteRT CLI를 구성한다.""" | |
| parser = argparse.ArgumentParser(description="Export Math Ink 0.6 P Formula student") | |
| parser.add_argument("--student-checkpoint", type=Path, required=True) | |
| parser.add_argument("--data", type=Path, required=True) | |
| parser.add_argument("--output", type=Path, required=True) | |
| parser.add_argument("--maximum-representative-samples", type=int, default=256) | |
| parser.add_argument("--convert-litert", action="store_true") | |
| return parser.parse_args() | |
| def main() -> None: | |
| """필요 변수: release gate 통과 student와 원 학습 corpus. 작동 원리: lineage 재검증 후 단일 formula graph의 parity를 고정한다.""" | |
| args = _parse_args() | |
| if args.maximum_representative_samples <= 0: | |
| raise ValueError("대표 입력 상한은 양수여야 합니다.") | |
| payload = torch.load( | |
| args.student_checkpoint, | |
| map_location="cpu", | |
| weights_only=False, | |
| ) | |
| data_sha256 = _file_sha25606(args.data) | |
| validate_p_formula_student_export06(payload, data_sha256=data_sha256) | |
| parent = args.student_checkpoint.parent | |
| base = _resolve_checkpoint_path06(payload["student_base_checkpoint"], parent=parent) | |
| online_checkpoint = _resolve_checkpoint_path06( | |
| payload["student_online_adapter"], | |
| parent=parent, | |
| ) | |
| engine = MathInk06Engine(base, adapter_checkpoint=online_checkpoint) | |
| formula_adapter = SkeletonTrajectoryAdapter06( | |
| hidden_size=int(payload["hidden_size"]), | |
| ) | |
| formula_adapter.load_state_dict(payload["state_dict"]) | |
| wrapper = PFormulaStudentExportWrapper06( | |
| engine.model, | |
| _online_branch06(engine.composite_adapter), | |
| formula_adapter, | |
| family_weight=engine.online_family_fusion_weight, | |
| exact_family_index=engine.exact_family_index, | |
| ).eval() | |
| labels = tuple(str(label) for label in engine.labels) | |
| representatives, audit = _representative_inputs06( | |
| args.data, | |
| labels=labels, | |
| maximum_samples=args.maximum_representative_samples, | |
| ) | |
| exported = torch.export.export(wrapper, representatives[0], strict=True) | |
| equivalence = exported_equivalence06(wrapper, exported, representatives) | |
| args.output.mkdir(parents=True, exist_ok=True) | |
| program_path = args.output / "p_formula_online.pt2" | |
| _save_exported_program06(exported, program_path) | |
| size_gate = program_path.stat().st_size <= MAXIMUM_MODEL_BYTES06 | |
| report: dict[str, Any] = { | |
| "schema": "aiflow-math-ink-06-p-formula-student-export-v1", | |
| "model_version": f"{engine.model_version}+p-formula-student", | |
| "exact_label_count": len(labels), | |
| "vocabulary_sha256": _vocabulary_sha25606(list(labels)), | |
| "student_checkpoint": str(args.student_checkpoint), | |
| "data_sha256": data_sha256, | |
| "teacher_seeds": [17, 31, 47], | |
| "teacher_weights_embedded": False, | |
| "preflight": audit, | |
| "representative_samples": len(representatives), | |
| "torch_version": torch.__version__, | |
| "torch_export": { | |
| **equivalence, | |
| "path": program_path.name, | |
| "bytes": program_path.stat().st_size, | |
| "maximum_model_bytes": MAXIMUM_MODEL_BYTES06, | |
| "size_gate_passed": size_gate, | |
| }, | |
| "torch_export_gate_passed": bool(equivalence["gate_passed"] and size_gate), | |
| "litert_package_available": importlib.util.find_spec("litert_torch") is not None, | |
| "litert": {"converted": False, "reason": "conversion_not_requested"}, | |
| "android_validation": False, | |
| "product_validation": False, | |
| } | |
| if args.convert_litert: | |
| if not report["litert_package_available"]: | |
| report["litert"] = { | |
| "converted": False, | |
| "reason": "litert_torch_not_installed", | |
| } | |
| else: | |
| report["litert"] = _convert_litert( | |
| wrapper, | |
| representatives, | |
| args.output / "p_formula_online.tflite", | |
| ) | |
| report["next_gate"] = ( | |
| "LiteRT top-1 100%·max logit error≤0.02, then Android low/mid/high tier benchmark" | |
| ) | |
| (args.output / "export_manifest.json").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() | |