"""Export Zap: bf16 safetensors (the published weights) + ONNX (fp32, optional int8) with parity checks. Usage: python export.py MODEL_DIR DATA_DIR OUT_DIR """ import json import os import shutil import sys import numpy as np import torch sys.path.insert(0, "/job/work") from train.zapdata import N_FORM, windows # noqa: E402 from train.zaptrain import pad, read_jsonl # noqa: E402 def main(model_dir, data_dir, out): from transformers import AutoTokenizer, BertForTokenClassification os.makedirs(os.path.join(out, "onnx"), exist_ok=True) tok = AutoTokenizer.from_pretrained(model_dir) m32 = BertForTokenClassification.from_pretrained(model_dir, dtype=torch.float32).eval() m_exp = BertForTokenClassification.from_pretrained(model_dir, dtype=torch.float32).eval() # tracing mutates it # 1) bf16 weights m16 = BertForTokenClassification.from_pretrained(model_dir, dtype=torch.bfloat16).eval() m16.save_pretrained(out) tok.save_pretrained(out) # portable tokenizer class: loads tokenizer.json as-is in transformers 4.x/5.x and transformers.js tc_path = os.path.join(out, "tokenizer_config.json") tc = json.load(open(tc_path)) tc["tokenizer_class"] = "PreTrainedTokenizerFast" json.dump(tc, open(tc_path, "w"), indent=1) tok2 = AutoTokenizer.from_pretrained(out) probe = ["[TITLE] Anmeldung – Café Müller [FLD] [TYPE] input password [NAME] passwort [LBL] Kennwort", "[TITLE] 登录 - 账户 [FLD] [TYPE] input text [PH] 手机号/邮箱 [FLD] [TYPE] input password", "[TITLE] Вход [FLD] [TYPE] input email [LBL] Электронная почта [REQ]"] for t in probe: assert tok(t)["input_ids"] == tok2(t)["input_ids"], "tokenizer round-trip mismatch" print("tokenizer round-trip ok", tok2.__class__.__name__) # 2) ONNX fp32 ids = torch.tensor([[tok.cls_token_id] + [5] * 30 + [tok.sep_token_id]] * 2) att = torch.ones_like(ids) class Wrap(torch.nn.Module): def __init__(self, m): super().__init__() self.m = m def forward(self, input_ids, attention_mask): return self.m(input_ids=input_ids, attention_mask=attention_mask).logits onnx_path = os.path.join(out, "onnx", "model.onnx") torch.onnx.export(Wrap(m_exp), (ids, att), onnx_path, input_names=["input_ids", "attention_mask"], output_names=["logits"], opset_version=17, dynamo=False, dynamic_axes={"input_ids": {0: "batch", 1: "seq"}, "attention_mask": {0: "batch", 1: "seq"}, "logits": {0: "batch", 1: "seq"}}) import onnxruntime as ort from onnxruntime.quantization import QuantType, quantize_dynamic q_path = os.path.join(out, "onnx", "model_quantized.onnx") quantize_dynamic(onnx_path, q_path, weight_type=QuantType.QInt8) sess = {"fp32": ort.InferenceSession(onnx_path, providers=["CPUExecutionProvider"]), "int8": ort.InferenceSession(q_path, providers=["CPUExecutionProvider"])} # 3) parity on test windows: argmax agreement at [CLS] / [FLD] positions vs torch fp32 rows = read_jsonl(os.path.join(data_dir, "test.jsonl"))[:3000] items = [w for r in rows for w in windows(tok, r["ctx"], r["fields"])] agree = {k: [0, 0] for k in ("bf16_torch", "fp32_onnx", "int8_onnx")} maxdiff = 0.0 for s in range(0, len(items), 64): chunk = items[s:s + 64] x = pad([c[0] for c in chunk], tok.pad_token_id) a = (x != tok.pad_token_id).long() with torch.no_grad(): ref = m32(input_ids=x, attention_mask=a).logits.float() b16 = m16(input_ids=x, attention_mask=a).logits.float() outs = {"bf16_torch": b16} for k, ss in sess.items(): outs[f"{k}_onnx"] = torch.from_numpy(ss.run(None, {"input_ids": x.numpy().astype(np.int64), "attention_mask": a.numpy().astype(np.int64)})[0]) maxdiff = max(maxdiff, float((outs["fp32_onnx"] - ref).abs().max())) for b, (_, pos, _) in enumerate(chunk): for p, sl in [(0, slice(0, N_FORM))] + [(p, slice(N_FORM, None)) for p in pos]: r = int(ref[b, p, sl].argmax()) for k, o in outs.items(): agree[k][0] += int(o[b, p, sl].argmax()) == r agree[k][1] += 1 report = {"positions": agree["fp32_onnx"][1], "fp32_onnx_max_abs_logit_diff": maxdiff, **{f"{k}_argmax_agreement": v[0] / max(1, v[1]) for k, v in agree.items()}} print(json.dumps(report, indent=1)) json.dump(report, open(os.path.join(out, "export_parity.json"), "w"), indent=1) if report["int8_onnx_argmax_agreement"] < 0.995: os.remove(q_path) print("int8 ONNX dropped: agreement below 99.5%") shutil.copy(os.path.join(model_dir, "val_metrics.json"), out) if __name__ == "__main__": main(*sys.argv[1:4])