File size: 8,007 Bytes
9d03fa1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
"""
verify.py — Validate the exported ONNX files against the original PyTorch model.

Runs five checks and prints a readable report:

  0. Artifact consistency — the shipped tokenizer/config vs the pinned source.
  1. Raw-logit parity — PyTorch export wrapper vs ONNX (fp32). Confirms the
     export itself is faithful (should be ~1e-4 or smaller).
  2. End-to-end parity — the full post-processed output (task types +
     complexity scores) reproduced from the ONNX logits via the model's own
     post-processing. Should match to the rounding the post-processing applies.
  3. Ground-truth anchor — the README example must classify as "Code Generation"
     with the documented complexity score.
  4. fp16 drift — fp16 outputs vs fp32; expected to be negligible (~1e-3).

Exit code is non-zero if any hard check fails.
"""

import os
import sys

import numpy as np
import onnxruntime as ort
import torch
from transformers import AutoConfig, AutoTokenizer

from export import (
    MODEL_NAME,
    MODEL_REVISION,
    OUT_DIR,
    OUTPUT_NAMES,
    ROOT_DIR,
    ExportWrapper,
    load_model,
)

# Diverse prompts spanning task types / complexity levels.
PROMPTS = [
    "Write a Python script that uses a for loop.",
    "What is the capital of France?",
    "Summarize the following report in three concise bullet points, keeping only "
    "the financial figures and omitting any commentary about strategy.",
    "Prove, step by step and with full rigor, that the square root of 2 is "
    "irrational, then explain where the argument would break for the square root of 4.",
]

# Documented output from README.md for the example prompt (reproduced verbatim,
# including the "Prompt: " prefix).
README_PROMPT = "Prompt: Write a Python script that uses a for loop."
README_EXPECTED_TASK_1 = "Code Generation"
README_EXPECTED_SCORE = 0.27823

# Numeric fields in the result dict (everything except the two string fields).
NUMERIC_FIELDS = [
    "task_type_prob",
    "creativity_scope",
    "reasoning",
    "contextual_knowledge",
    "number_of_few_shots",
    "domain_knowledge",
    "no_label_reason",
    "constraint_ct",
    "prompt_complexity_score",
]
STRING_FIELDS = ["task_type_1", "task_type_2"]


def encode(tok, prompt):
    return tok(prompt, return_tensors="pt", truncation=True, max_length=512)


def run_onnx(sess, enc):
    """Return the 8 raw logit arrays in OUTPUT_NAMES order."""
    return sess.run(
        None,
        {
            "input_ids": enc["input_ids"].numpy(),
            "attention_mask": enc["attention_mask"].numpy(),
        },
    )


def create_session(path):
    """Create a quiet, deterministic CPU session for release validation."""
    options = ort.SessionOptions()
    options.log_severity_level = 3
    return ort.InferenceSession(
        path,
        sess_options=options,
        providers=["CPUExecutionProvider"],
    )


def result_from_onnx(model, onnx_logits):
    """Reuse the model's own post-processing on ONNX logits -> result dict."""
    return model.process_logits([torch.tensor(x) for x in onnx_logits])


def dict_diff(a, b):
    """Max abs numeric drift and any string mismatch between two result dicts."""
    max_num = 0.0
    string_mismatch = None
    for f in NUMERIC_FIELDS:
        av = np.array(a[f], dtype=float)
        bv = np.array(b[f], dtype=float)
        max_num = max(max_num, float(np.abs(av - bv).max()))
    for f in STRING_FIELDS:
        if a[f] != b[f]:
            string_mismatch = (f, a[f], b[f])
    return max_num, string_mismatch


def main():
    ok = True
    print("== Check 0: shipped tokenizer/config consistency ==")
    tok = AutoTokenizer.from_pretrained(ROOT_DIR, local_files_only=True)
    source_tok = AutoTokenizer.from_pretrained(
        MODEL_NAME,
        revision=MODEL_REVISION,
    )
    local_config = AutoConfig.from_pretrained(ROOT_DIR, local_files_only=True)
    source_config = AutoConfig.from_pretrained(
        MODEL_NAME,
        revision=MODEL_REVISION,
    )

    tokenizer_matches = all(
        encode(tok, p)["input_ids"].equal(encode(source_tok, p)["input_ids"])
        and encode(tok, p)["attention_mask"].equal(
            encode(source_tok, p)["attention_mask"]
        )
        for p in PROMPTS + [README_PROMPT]
    )
    config_fields = ["target_sizes", "task_type_map", "weights_map", "divisor_map"]
    config_matches = all(
        getattr(local_config, field) == getattr(source_config, field)
        for field in config_fields
    )
    print(f"  tokenizer matches pinned source: {tokenizer_matches}")
    print(f"  scoring config matches pinned source: {config_matches}")
    if not tokenizer_matches or not config_matches:
        print("  [FAIL] shipped preprocessing artifacts differ from pinned source")
        ok = False
    else:
        print("  [ok]")

    print("Loading PyTorch model ...")
    model = load_model()
    wrapper = ExportWrapper(model).eval()

    fp32 = create_session(os.path.join(OUT_DIR, "model.onnx"))
    fp16 = create_session(os.path.join(OUT_DIR, "model_fp16.onnx"))

    for name, session in [("fp32", fp32), ("fp16", fp16)]:
        output_names = [output.name for output in session.get_outputs()]
        if output_names != OUTPUT_NAMES:
            print(
                f"  [FAIL] {name} output order is {output_names}, "
                f"expected {OUTPUT_NAMES}"
            )
            ok = False

    print("\n== Check 1: raw-logit parity (PyTorch vs ONNX fp32) ==")
    max_logit_diff = 0.0
    for p in PROMPTS:
        enc = encode(tok, p)
        with torch.no_grad():
            pt_logits = wrapper(enc["input_ids"], enc["attention_mask"])
        onnx_logits = run_onnx(fp32, enc)
        for a, b in zip(pt_logits, onnx_logits):
            max_logit_diff = max(max_logit_diff, float(np.abs(a.numpy() - b).max()))
    print(f"  max |logit_pt - logit_onnx| = {max_logit_diff:.2e}")
    if max_logit_diff > 1e-3:
        print("  [FAIL] logit drift larger than 1e-3")
        ok = False
    else:
        print("  [ok]")

    print("\n== Check 2: end-to-end parity (PyTorch vs ONNX-derived) ==")
    e2e_max = 0.0
    for p in PROMPTS:
        enc = encode(tok, p)
        ref = model(enc)
        got = result_from_onnx(model, run_onnx(fp32, enc))
        num, mism = dict_diff(ref, got)
        e2e_max = max(e2e_max, num)
        if mism:
            print(f"  [FAIL] string mismatch on {p[:40]!r}: {mism}")
            ok = False
    print(f"  max numeric drift = {e2e_max:.2e}")
    print("  [ok]" if e2e_max <= 1e-3 else "  [FAIL] end-to-end drift > 1e-3")
    ok = ok and e2e_max <= 1e-3

    print("\n== Check 3: README ground-truth anchor ==")
    enc = encode(tok, README_PROMPT)
    ref = result_from_onnx(model, run_onnx(fp32, enc))
    got_task = ref["task_type_1"][0]
    got_score = ref["prompt_complexity_score"][0]
    print(f"  task_type_1 = {got_task!r} (expected {README_EXPECTED_TASK_1!r})")
    print(f"  prompt_complexity_score = {got_score} (expected ~{README_EXPECTED_SCORE})")
    if got_task != README_EXPECTED_TASK_1 or abs(got_score - README_EXPECTED_SCORE) > 1e-3:
        print("  [FAIL] does not match documented output")
        ok = False
    else:
        print("  [ok]")

    print("\n== Check 4: fp16 drift (fp16 vs fp32) ==")
    fp16_max = 0.0
    for p in PROMPTS:
        enc = encode(tok, p)
        ref = result_from_onnx(model, run_onnx(fp32, enc))
        got = result_from_onnx(model, run_onnx(fp16, enc))
        num, mism = dict_diff(ref, got)
        fp16_max = max(fp16_max, num)
        if mism:
            print(f"  [FAIL] fp16 changed a task label on {p[:40]!r}: {mism}")
            ok = False
    print(f"  max numeric drift (fp16 vs fp32) = {fp16_max:.2e}")
    if fp16_max > 1e-2:
        print("  [FAIL] fp16 drift larger than 1e-2")
        ok = False
    else:
        print("  [ok]")

    print("\n" + ("ALL HARD CHECKS PASSED" if ok else "SOME CHECKS FAILED"))
    sys.exit(0 if ok else 1)


if __name__ == "__main__":
    main()