aiflow-math-ink-06-intermediate / scripts /export_math_ink_06_p_formula_student.py
cwLeeDev's picture
Separate public research conversion from private P mobile release
5ee3a5e verified
Raw
History Blame Contribute Delete
10.7 kB
"""통과한 P Formula student를 strict torch.export와 선택적 LiteRT로 고정한다."""
from __future__ import annotations
import argparse
import gzip
from hashlib import sha256
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.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,
)
MAXIMUM_MODEL_BYTES06 = 25 * 1024 * 1024
def _file_sha25606(path: Path) -> str:
"""필요 변수: P corpus. 작동 원리: trainer import 없이 원본 byte-level SHA-256을 계산한다."""
digest = sha256()
with path.open("rb") as stream:
for chunk in iter(lambda: stream.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _read_jsonl06(path: Path) -> list[dict[str, Any]]:
"""필요 변수: UTF-8 JSONL 또는 gzip JSONL. 작동 원리: 제품 export에 필요한 record만 독립적으로 읽는다."""
stream = (
gzip.open(path, "rt", encoding="utf-8")
if path.suffix == ".gz"
else path.open("r", encoding="utf-8")
)
with stream:
return [
json.loads(line)
for line in stream
if line.strip()
]
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이 모두 필요합니다.")
for field in ("student_base_sha256", "student_online_adapter_sha256"):
value = str(payload.get(field) or "")
if len(value) != 64 or any(character not in "0123456789abcdef" for character in value):
raise ValueError(f"Student lineage의 {field}가 유효하지 않습니다.")
def validate_p_formula_student_artifacts06(
payload: dict[str, Any],
*,
base_checkpoint: Path,
online_adapter: Path,
) -> None:
"""필요 변수: student metadata·실제 base/adapter. 작동 원리: 경로가 아닌 byte hash로 가중치 lineage를 검증한다."""
expected = {
"student_base_sha256": _file_sha25606(base_checkpoint),
"student_online_adapter_sha256": _file_sha25606(online_adapter),
}
for field, actual in expected.items():
if str(payload.get(field) or "") != actual:
raise ValueError(f"Student lineage와 실제 {field} artifact가 다릅니다.")
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 = _read_jsonl06(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,
)
validate_p_formula_student_artifacts06(
payload,
base_checkpoint=base,
online_adapter=online_checkpoint,
)
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()