File size: 7,696 Bytes
c228b1d 5ee3a5e c228b1d 5ee3a5e c228b1d 5ee3a5e c228b1d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 | """통과한 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()
|