aiflow-math-ink-06-intermediate / scripts /export_math_ink_06_p_mobile_pair.py
cwLeeDev's picture
Separate public research conversion from private P mobile release
5ee3a5e verified
Raw
History Blame Contribute Delete
9.63 kB
"""통과한 P Formula student online과 5-output raster를 하나의 모바일 쌍으로 export한다."""
from __future__ import annotations
import argparse
import importlib.util
import json
from pathlib import Path
import sys
from typing import Any
import numpy as np
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_canonical import render_canonical_ink
from math_grid_drawer.research.ink06_export import (
PFormulaStudentExportWrapper06,
RasterDebugExportWrapper06,
exported_equivalence06,
)
from math_grid_drawer.research.math_ink_06 import MathInk06Engine
from math_grid_drawer.research.p_formula_dataset06 import (
_formula_box06,
p_formula_symbol_ink06,
)
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.export_math_ink_06_p_formula_student import (
_online_branch06,
_file_sha25606,
_read_jsonl06,
_resolve_checkpoint_path06,
validate_p_formula_student_artifacts06,
validate_p_formula_student_export06,
)
MAXIMUM_MODEL_BUNDLE_BYTES06 = 25 * 1024 * 1024
def paired_p_representatives06(
records: list[dict[str, Any]],
*,
maximum_samples: int,
) -> tuple[list[tuple[torch.Tensor, ...]], list[tuple[torch.Tensor, ...]]]:
"""필요 변수: P test formula·상한. 작동 원리: 같은 symbol에서 online 128×19와 raster 128×128을 함께 만든다."""
if maximum_samples <= 0:
raise ValueError("대표 입력 상한은 양수여야 합니다.")
online: list[tuple[torch.Tensor, ...]] = []
raster: list[tuple[torch.Tensor, ...]] = []
for record in records:
if str(record.get("split") or "") != "test":
continue
formula_box = _formula_box06(record)
for symbol in record["symbols"]:
ink = p_formula_symbol_ink06(symbol, formula_box=formula_box)
image = np.asarray(render_canonical_ink(ink), dtype=np.float32)
online.append((torch.from_numpy(ink.features).unsqueeze(0),))
raster.append((
torch.from_numpy(1.0 - image / 255.0).unsqueeze(0).unsqueeze(0),
))
if len(online) >= maximum_samples:
return online, raster
if not online:
raise ValueError("P Formula test representative가 없습니다.")
return online, raster
def _raster_branch06(adapter: torch.nn.Module) -> torch.nn.Module:
"""필요 변수: single/dual adapter. 작동 원리: image virtual stroke에 대응하는 raster branch만 반환한다."""
return adapter.raster if isinstance(adapter, DualModalityTrajectoryAdapter06) else adapter
def main() -> None:
"""필요 변수: 통과 student·동일 P corpus·출력. 작동 원리: 같은 lineage의 online/raster graph를 export한다."""
parser = argparse.ArgumentParser(description="Export Math Ink 0.6 P mobile pair")
parser.add_argument("--student-checkpoint", type=Path, required=True)
parser.add_argument("--base-checkpoint", type=Path)
parser.add_argument("--adapter-checkpoint", type=Path)
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")
args = parser.parse_args()
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 = args.base_checkpoint or _resolve_checkpoint_path06(
payload["student_base_checkpoint"],
parent=parent,
)
adapter_checkpoint = args.adapter_checkpoint or _resolve_checkpoint_path06(
payload["student_online_adapter"],
parent=parent,
)
validate_p_formula_student_artifacts06(
payload,
base_checkpoint=base,
online_adapter=adapter_checkpoint,
)
engine = MathInk06Engine(base, adapter_checkpoint=adapter_checkpoint)
formula_adapter = SkeletonTrajectoryAdapter06(
hidden_size=int(payload["hidden_size"]),
)
formula_adapter.load_state_dict(payload["state_dict"])
online_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()
fusion = engine.raster_fusion
if any(float(fusion[key]) != 0.0 for key in ("family_weight", "geometry_weight", "symmetry_weight")):
raise ValueError("Mobile raster export는 학습 graph 밖 auxiliary fusion을 허용하지 않습니다.")
raster_wrapper = RasterDebugExportWrapper06(
engine.model,
adapter=_raster_branch06(engine.composite_adapter),
fusion_mode=str(fusion["mode"]),
score_weight=float(fusion["score_weight"]),
).eval()
records = _read_jsonl06(args.data)
audit = audit_p_formula_records06(records)
if not audit["eligible_for_product_evaluation"]:
raise ValueError("P Formula corpus가 product preflight를 통과하지 못했습니다.")
online_inputs, raster_inputs = paired_p_representatives06(
records,
maximum_samples=args.maximum_representative_samples,
)
online_export = torch.export.export(
online_wrapper,
online_inputs[0],
strict=True,
)
raster_export = torch.export.export(
raster_wrapper,
raster_inputs[0],
strict=True,
)
equivalence = {
"online": exported_equivalence06(
online_wrapper,
online_export,
online_inputs,
),
"raster": exported_equivalence06(
raster_wrapper,
raster_export,
raster_inputs,
),
}
args.output.mkdir(parents=True, exist_ok=True)
online_path = args.output / "p_formula_online.pt2"
raster_path = args.output / "raster_debug5.pt2"
_save_exported_program06(online_export, online_path)
_save_exported_program06(raster_export, raster_path)
total_bytes = online_path.stat().st_size + raster_path.stat().st_size
size_gate = total_bytes <= MAXIMUM_MODEL_BUNDLE_BYTES06
labels = tuple(str(label) for label in engine.labels)
report: dict[str, Any] = {
"schema": "aiflow-math-ink-06-p-mobile-pair-export-v1",
"model_version": f"{engine.model_version}+p-formula-student",
"student_checkpoint": str(args.student_checkpoint),
"data_sha256": data_sha256,
"teacher_seeds": [17, 31, 47],
"teacher_weights_embedded": False,
"exact_label_count": len(labels),
"vocabulary_sha256": _vocabulary_sha25606(list(labels)),
"raster_output_count": 5,
"representative_samples": len(online_inputs),
"preflight": audit,
"torch_version": torch.__version__,
"torch_export": {
"online": {
**equivalence["online"],
"path": online_path.name,
"bytes": online_path.stat().st_size,
},
"raster": {
**equivalence["raster"],
"path": raster_path.name,
"bytes": raster_path.stat().st_size,
},
"total_bytes": total_bytes,
"maximum_bundle_bytes": MAXIMUM_MODEL_BUNDLE_BYTES06,
"size_gate_passed": size_gate,
},
"torch_export_gate_passed": bool(
size_gate
and equivalence["online"]["gate_passed"]
and equivalence["raster"]["gate_passed"]
),
"litert_package_available": importlib.util.find_spec("litert_torch") is not None,
"litert": {
"online": {"converted": False, "reason": "conversion_not_requested"},
"raster": {"converted": False, "reason": "conversion_not_requested"},
},
"product_validation": False,
}
if args.convert_litert:
if not report["litert_package_available"]:
for branch in ("online", "raster"):
report["litert"][branch] = {
"converted": False,
"reason": "litert_torch_not_installed",
}
else:
report["litert"] = {
"online": _convert_litert(
online_wrapper,
online_inputs,
args.output / "p_formula_online.tflite",
),
"raster": _convert_litert(
raster_wrapper,
raster_inputs,
args.output / "raster_debug5.tflite",
),
}
report["next_gate"] = (
"package exact model pair, raster release validation, Android low/mid/high"
)
(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()