"""Math Ink 0.6을 strict torch.export로 고정하고 선택적으로 LiteRT로 변환한다.""" from __future__ import annotations import argparse from hashlib import sha256 import importlib.util import json from pathlib import Path import sys import numpy as np import torch PROJECT_ROOT = Path(__file__).parents[1] SOURCE_ROOT = PROJECT_ROOT / "src" if str(SOURCE_ROOT) not in sys.path: sys.path.insert(0, str(SOURCE_ROOT)) from math_grid_drawer.research.ink06_canonical import canonicalize_ink06, render_canonical_ink from math_grid_drawer.research.ink06_export import ( OnlineExportWrapper06, RasterDebugExportWrapper06, exported_equivalence06, ) from math_grid_drawer.research.math_ink_06 import MathInk06Engine from math_grid_drawer.research.skeleton_adapter06 import DualModalityTrajectoryAdapter06 def _vocabulary_sha25606(labels: tuple[str, ...] | list[str]) -> str: """필요 변수: 순서가 고정된 exact labels. 작동 원리: Android label table과 graph의 동일성을 위한 SHA-256을 만든다.""" payload = json.dumps( list(labels), ensure_ascii=False, separators=(",", ":"), ).encode("utf-8") return sha256(payload).hexdigest() def _representative_inputs(baseline_report: Path, data_path: Path) -> tuple[list[tuple[torch.Tensor, ...]], list[tuple[torch.Tensor, ...]]]: """필요 변수: strict baseline·HWRT JSONL. 작동 원리: 고정 76개를 128×19와 128×128 대표 입력으로 재구성한다.""" from math_grid_drawer.research.external_corpus import read_jsonl baseline = json.loads(baseline_report.read_text(encoding="utf-8")) accepted_ids = {row["sample_id"] for row in baseline["rows"] if row["raster_gate"]} records = {row["sample_id"]: row for row in read_jsonl(data_path) if row["sample_id"] in accepted_ids} online_inputs: list[tuple[torch.Tensor, ...]] = [] raster_inputs: list[tuple[torch.Tensor, ...]] = [] for row in baseline["rows"]: if not row["raster_gate"]: continue record = records[row["sample_id"]] ink = canonicalize_ink06(record["strokes"], canvas_width=768, canvas_height=128, trust_timestamps=False) online_inputs.append((torch.from_numpy(ink.features).unsqueeze(0),)) image = np.asarray(render_canonical_ink(ink), dtype=np.float32) raster_inputs.append((torch.from_numpy(1.0 - image / 255.0).unsqueeze(0).unsqueeze(0),)) return online_inputs, raster_inputs def _load_representative_inputs06( cache_path: Path, ) -> tuple[list[tuple[torch.Tensor, ...]], list[tuple[torch.Tensor, ...]]]: """필요 변수: 고정 representative cache. 작동 원리: Colab에서도 같은 76개 batch-1 입력을 복원한다.""" payload = torch.load(cache_path, map_location="cpu", weights_only=True) if payload.get("schema") != "aiflow-math-ink-06-export-inputs-v1": raise ValueError("지원하지 않는 export representative cache입니다.") online, raster = payload["online"], payload["raster"] if online.ndim != 3 or online.shape[1:] != (128, 19): raise ValueError(f"online representative shape가 다릅니다: {tuple(online.shape)}") if raster.ndim != 4 or raster.shape[1:] != (1, 128, 128): raise ValueError(f"raster representative shape가 다릅니다: {tuple(raster.shape)}") if len(online) != len(raster) or not len(online): raise ValueError("online/raster representative 분모가 다릅니다.") return ( [(value.unsqueeze(0),) for value in online], [(value.unsqueeze(0),) for value in raster], ) def _save_representative_inputs06( output: Path, online_inputs: list[tuple[torch.Tensor, ...]], raster_inputs: list[tuple[torch.Tensor, ...]], ) -> None: """필요 변수: 두 대표 입력 목록·출력. 작동 원리: 실제 HWRT-derived 입력만 tensor cache로 고정한다.""" torch.save({ "schema": "aiflow-math-ink-06-export-inputs-v1", "online": torch.cat([row[0] for row in online_inputs], dim=0).cpu(), "raster": torch.cat([row[0] for row in raster_inputs], dim=0).cpu(), "samples": len(online_inputs), "contains_labels": False, "product_validation": False, }, output) def _convert_litert( wrapper: torch.nn.Module, samples: list[tuple[torch.Tensor, ...]], output: Path, ) -> dict: """필요 변수: export 호환 wrapper·대표 입력 전체·출력. 작동 원리: 공식 converter 뒤 76개 top-1/logit parity를 검사한다.""" import litert_torch # type: ignore[import-not-found] edge_model = litert_torch.convert(wrapper.eval(), samples[0]) top1_matches = 0 max_error = 0.0 with torch.inference_mode(): for sample in samples: eager = wrapper(*sample) edge = edge_model(*sample) eager_values = eager if isinstance(eager, tuple) else (eager,) edge_values = edge if isinstance(edge, tuple) else (edge,) if len(eager_values) != len(edge_values): raise ValueError("PyTorch/LiteRT 출력 개수가 다릅니다.") max_error = max(max_error, max( float(np.max(np.abs(left.detach().cpu().numpy() - np.asarray(right)))) for left, right in zip(eager_values, edge_values) )) top1_matches += int( np.argmax(eager_values[0].detach().cpu().numpy(), axis=-1)[0] == np.argmax(np.asarray(edge_values[0]), axis=-1)[0] ) edge_model.export(str(output)) return { "converted": True, "path": output.name, "bytes": output.stat().st_size, "samples": len(samples), "top1_matches": top1_matches, "top1_agreement": top1_matches / max(len(samples), 1), "max_absolute_logit_error": max_error, "gate_passed": top1_matches == len(samples) and max_error <= 0.02, } def _load_composite06( checkpoint: Path, adapter_checkpoint: Path, ) -> tuple[MathInk06Engine, torch.nn.Module, dict]: """필요 변수: base·adapter checkpoint. 작동 원리: base→shared state→modality adapter 순서로 배포 모델을 합성한다.""" engine = MathInk06Engine(checkpoint, adapter_checkpoint=adapter_checkpoint) payload = torch.load(adapter_checkpoint, map_location="cpu", weights_only=False) return engine, engine.composite_adapter, payload def _adapter_branches06(adapter: torch.nn.Module) -> tuple[torch.nn.Module, torch.nn.Module]: """필요 변수: single/dual adapter. 작동 원리: export graph에서 데이터 의존 분기 없이 online/raster branch를 고정한다.""" if isinstance(adapter, DualModalityTrajectoryAdapter06): return adapter.online, adapter.raster return adapter, adapter def _save_exported_program06(exported: torch.export.ExportedProgram, output: Path) -> None: """필요 변수: export program·목표 파일. 작동 원리: stale ZIP 재사용 없이 임시 파일을 원자적으로 교체한다.""" temporary = output.with_suffix(output.suffix + ".part") if temporary.exists(): temporary.unlink() torch.export.save(exported, temporary) temporary.replace(output) def main() -> None: """필요 변수: checkpoint·strict 입력·출력·변환 선택. 작동 원리: 두 경로의 export/동등성/선택적 LiteRT 결과를 manifest로 고정한다.""" parser = argparse.ArgumentParser(description="Export Math Ink 0.6 for LiteRT") parser.add_argument("--checkpoint", type=Path, required=True) parser.add_argument("--adapter-checkpoint", type=Path, required=True) parser.add_argument("--representative-inputs", type=Path) parser.add_argument("--save-representative-inputs", type=Path) parser.add_argument("--baseline-report", type=Path, default=PROJECT_ROOT / "research/runs/full_model_stage2_20260722/isolated_checkpoint.json") parser.add_argument("--data", type=Path, default=PROJECT_ROOT / "research/data/open_pretrain/hwrt_expanded_v2/hwrt_expanded.jsonl.gz") parser.add_argument("--output", type=Path, required=True) parser.add_argument("--convert-litert", action="store_true") args = parser.parse_args() engine, adapter, adapter_payload = _load_composite06(args.checkpoint, args.adapter_checkpoint) fusion = engine.raster_fusion if any(float(fusion[key]) != 0.0 for key in ("family_weight", "geometry_weight", "symmetry_weight")): raise ValueError("현재 LiteRT wrapper는 family/geometry/symmetry 보조 fusion을 지원하지 않습니다.") online_adapter, raster_adapter = _adapter_branches06(adapter) online = OnlineExportWrapper06( engine.model, online_adapter, family_weight=engine.online_family_fusion_weight, exact_family_index=engine.exact_family_index, ).eval() raster = RasterDebugExportWrapper06( engine.model, adapter=raster_adapter, fusion_mode=str(fusion["mode"]), score_weight=float(fusion["score_weight"]), ).eval() if args.representative_inputs: online_inputs, raster_inputs = _load_representative_inputs06(args.representative_inputs) else: online_inputs, raster_inputs = _representative_inputs(args.baseline_report, args.data) if args.save_representative_inputs: args.save_representative_inputs.parent.mkdir(parents=True, exist_ok=True) _save_representative_inputs06( args.save_representative_inputs, online_inputs, raster_inputs, ) online_export = torch.export.export(online, online_inputs[0], strict=True) raster_export = torch.export.export(raster, raster_inputs[0], strict=True) args.output.mkdir(parents=True, exist_ok=True) online_path, raster_path = args.output / "online.pt2", args.output / "raster.pt2" _save_exported_program06(online_export, online_path) _save_exported_program06(raster_export, raster_path) report = { "schema": "aiflow-math-ink-06-dual-export-v1", "checkpoint": str(args.checkpoint), "adapter_checkpoint": str(args.adapter_checkpoint), "adapter_architecture": str(adapter_payload["adapter_architecture"]), "shared_state_applied": bool(adapter_payload.get("shared_state_dict")), "model_version": engine.model_version, "exact_label_count": len(engine.labels), "vocabulary_sha256": _vocabulary_sha25606(list(engine.labels)), "raster_output_count": 5, "online_family_fusion_weight": engine.online_family_fusion_weight, "torch_version": torch.__version__, "torch_export": { "online": {**exported_equivalence06(online, online_export, online_inputs), "path": online_path.name, "bytes": online_path.stat().st_size}, "raster": {**exported_equivalence06(raster, raster_export, raster_inputs), "path": raster_path.name, "bytes": raster_path.stat().st_size}, }, "litert_package_available": importlib.util.find_spec("litert_torch") is not None, "litert": {"converted": False, "reason": "conversion_not_requested"}, "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"] = { "online": _convert_litert(online, online_inputs, args.output / "online.tflite"), "raster": _convert_litert(raster, raster_inputs, args.output / "raster.tflite"), } report["torch_export_gate_passed"] = all( bool(report["torch_export"][name]["gate_passed"]) for name in ("online", "raster") ) (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()