| """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 |
|
|
| 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() |
|
|