aiflow-math-ink-06-intermediate / scripts /package_math_ink_06_mobile_models.py
cwLeeDev's picture
Separate public research conversion from private P mobile release
5ee3a5e verified
Raw
History Blame Contribute Delete
7.7 kB
"""통과한 online/raster LiteRT를 하나의 Android 모델 쌍으로 묶는다."""
from __future__ import annotations
import argparse
from datetime import datetime, timezone
from hashlib import sha256
import json
from pathlib import Path
import re
from typing import Any
MAXIMUM_BUNDLE_BYTES06 = 25 * 1024 * 1024
def _file_sha256_06(path: Path) -> str:
"""필요 변수: 모델 파일. 작동 원리: 파일 전체를 streaming 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 model_bundle_sha25606(online_sha256: str, raster_sha256: str) -> str:
"""필요 변수: online/raster hash. 작동 원리: Android와 동일한 ordered bundle 지문을 만든다."""
payload = f"online:{online_sha256}\nraster:{raster_sha256}\n".encode("utf-8")
return sha256(payload).hexdigest()
def _require_litert_row06(
report: dict[str, Any],
*,
branch: str,
) -> dict[str, Any]:
"""필요 변수: export report·branch. 작동 원리: 변환·parity gate가 모두 통과한 LiteRT 행만 반환한다."""
litert = report.get("litert") or {}
row = litert if branch == "online" and "online" not in litert else litert.get(branch)
if not isinstance(row, dict):
raise ValueError(f"{branch} LiteRT 결과가 없습니다.")
if row.get("converted") is not True or row.get("gate_passed") is not True:
raise ValueError(f"{branch} LiteRT 변환/parity gate가 통과하지 않았습니다.")
if float(row.get("top1_agreement", 0.0)) != 1.0:
raise ValueError(f"{branch} LiteRT top-1 agreement가 100%가 아닙니다.")
if float(row.get("max_absolute_logit_error", float("inf"))) > 0.02:
raise ValueError(f"{branch} LiteRT logit 오차가 0.02를 초과했습니다.")
return row
def package_mobile_models06(
*,
online_report: dict[str, Any],
raster_report: dict[str, Any],
online_model: Path,
raster_model: Path,
) -> dict[str, Any]:
"""필요 변수: 두 export report와 실제 flatbuffer. 작동 원리: vocabulary·parity·파일을 검증해 불변 모델 쌍을 만든다."""
online_schema = str(online_report.get("schema") or "")
raster_schema = str(raster_report.get("schema") or "")
pair_schema = "aiflow-math-ink-06-p-mobile-pair-export-v1"
if online_schema not in {
"aiflow-math-ink-06-p-formula-student-export-v1",
pair_schema,
}:
raise ValueError("Online은 통과한 P Formula student export여야 합니다.")
if raster_schema not in {"aiflow-math-ink-06-dual-export-v1", pair_schema}:
raise ValueError("Raster는 0.6 dual export여야 합니다.")
if pair_schema in {online_schema, raster_schema}:
if online_schema != pair_schema or raster_schema != pair_schema:
raise ValueError("P mobile pair report는 online/raster 양쪽에 함께 사용해야 합니다.")
lineage_keys = (
"student_checkpoint",
"data_sha256",
"model_version",
"vocabulary_sha256",
)
if any(
online_report.get(key) != raster_report.get(key)
for key in lineage_keys
):
raise ValueError("Online/raster P mobile pair lineage가 다릅니다.")
if online_report.get("torch_export_gate_passed") is not True:
raise ValueError("Online torch.export gate가 통과하지 않았습니다.")
if raster_report.get("torch_export_gate_passed") is not True:
raise ValueError("Raster torch.export gate가 통과하지 않았습니다.")
if int(raster_report.get("raster_output_count", 0)) != 5:
raise ValueError("Raster graph는 top-4 debug를 포함한 5-output이어야 합니다.")
counts = {
int(online_report.get("exact_label_count", 0)),
int(raster_report.get("exact_label_count", 0)),
}
vocabularies = {
str(online_report.get("vocabulary_sha256") or ""),
str(raster_report.get("vocabulary_sha256") or ""),
}
if counts != {378}:
raise ValueError("Online/raster 모두 378 exact labels여야 합니다.")
if (
len(vocabularies) != 1
or re.fullmatch(r"[0-9a-f]{64}", next(iter(vocabularies))) is None
):
raise ValueError("Online/raster vocabulary SHA-256이 같아야 합니다.")
online_row = _require_litert_row06(online_report, branch="online")
raster_row = _require_litert_row06(raster_report, branch="raster")
artifacts = {}
for name, path, row in (
("online", online_model, online_row),
("raster", raster_model, raster_row),
):
if not path.is_file():
raise FileNotFoundError(f"{name} LiteRT 파일이 없습니다: {path}")
size = path.stat().st_size
if Path(str(row.get("path") or "")).name != path.name:
raise ValueError(f"{name} report path와 실제 파일명이 다릅니다.")
if int(row.get("bytes", -1)) != size:
raise ValueError(f"{name} report byte 수와 실제 파일이 다릅니다.")
artifacts[name] = {
"path": path.name,
"bytes": size,
"sha256": _file_sha256_06(path),
}
total_bytes = sum(row["bytes"] for row in artifacts.values())
size_gate = total_bytes <= MAXIMUM_BUNDLE_BYTES06
bundle_hash = model_bundle_sha25606(
artifacts["online"]["sha256"],
artifacts["raster"]["sha256"],
)
return {
"schema": "aiflow-math-ink-06-mobile-model-bundle-v1",
"generated_at": datetime.now(timezone.utc).isoformat(),
"model_version": str(online_report["model_version"]),
"data_sha256": str(online_report.get("data_sha256") or ""),
"exact_label_count": 378,
"vocabulary_sha256": next(iter(vocabularies)),
"artifacts": artifacts,
"model_bundle_sha256": bundle_hash,
"total_bytes": total_bytes,
"maximum_bundle_bytes": MAXIMUM_BUNDLE_BYTES06,
"checks": {
"online_litert_parity": True,
"raster_litert_parity": True,
"raster_five_outputs": True,
"same_vocabulary": True,
"size": size_gate,
},
"package_gate_passed": size_gate,
"product_validation": False,
}
def main() -> None:
"""필요 변수: report·flatbuffer·출력. 작동 원리: 검증된 UTF-8 bundle manifest를 원자적으로 기록한다."""
parser = argparse.ArgumentParser(description="Package Math Ink 0.6 mobile models")
parser.add_argument("--online-report", type=Path, required=True)
parser.add_argument("--raster-report", type=Path, required=True)
parser.add_argument("--online-model", type=Path, required=True)
parser.add_argument("--raster-model", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
args = parser.parse_args()
result = package_mobile_models06(
online_report=json.loads(args.online_report.read_text(encoding="utf-8")),
raster_report=json.loads(args.raster_report.read_text(encoding="utf-8")),
online_model=args.online_model,
raster_model=args.raster_model,
)
args.output.parent.mkdir(parents=True, exist_ok=True)
temporary = args.output.with_suffix(args.output.suffix + ".part")
temporary.write_text(
json.dumps(result, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
temporary.replace(args.output)
print(json.dumps(result, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()