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