aiflow-math-ink-06-intermediate / scripts /export_math_ink_06_litert.py
cwLeeDev's picture
Expose raster top-4 strokes and add end-to-end release lineage gate
c228b1d verified
Raw
History Blame Contribute Delete
11.9 kB
"""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()