andreekeberg's picture
Add ONNX conversion
9d03fa1 unverified
Raw
History Blame Contribute Delete
8.01 kB
"""
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()