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