Zap / training /train /export.py
ProCreations's picture
Zap v1: 20.7M-param web form & field classifier for autofill (bf16, MIT)
e4805d0 verified
Raw History Blame Contribute Delete
4.97 kB
"""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])