CoderBak's picture
Publish attributed EditLens model kit: FP32 default, FP16, experimental INT8
f7cb4b0 verified
Raw
History Blame Contribute Delete
5.89 kB
"""Reproducible conversions of the pinned Pangram EditLens checkpoint.
License: CC-BY-NC-SA-4.0. See LICENSE and NOTICE in the repository root.
Run each stage in a separate process to bound peak conversion memory.
"""
from __future__ import annotations
import argparse
import hashlib
import json
from pathlib import Path
UPSTREAM = "pangram/editlens_roberta-large"
REVISION = "f93e1ace74528cfb48f337ab2fe946fb71a728cb"
def digest(path: Path) -> str:
with path.open("rb") as f:
return hashlib.file_digest(f, "sha256").hexdigest()
def verify_source(source: Path, output: Path) -> None:
metadata = json.loads((output / "upstream/metadata.json").read_text())
if metadata["repo_id"] != UPSTREAM or metadata["revision"] != REVISION:
raise RuntimeError("Unexpected upstream identity")
for info in metadata["files"]:
if info["name"] in {"README.md", ".gitattributes"}:
continue
if digest(source / info["name"]) != info["sha256"]:
raise RuntimeError("Source differs from pinned upstream: " + info["name"])
def export(source: Path, output: Path) -> None:
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
torch.set_num_threads(4)
model = AutoModelForSequenceClassification.from_pretrained(
source, local_files_only=True, dtype=torch.float32,
attn_implementation="eager",
).eval()
tokenizer = AutoTokenizer.from_pretrained(source, local_files_only=True)
class Classifier(torch.nn.Module):
def __init__(self, wrapped):
super().__init__()
self.wrapped = wrapped
def forward(self, input_ids, attention_mask):
return self.wrapped(input_ids=input_ids, attention_mask=attention_mask).logits
sample = tokenizer("A short example used only to trace the classifier graph.", return_tensors="pt")
with torch.inference_mode():
torch.onnx.export(
Classifier(model), (sample["input_ids"], sample["attention_mask"]),
str(output / "onnx/model.onnx"),
input_names=["input_ids", "attention_mask"], output_names=["logits"],
dynamic_axes={"input_ids": {0: "batch", 1: "sequence"},
"attention_mask": {0: "batch", 1: "sequence"},
"logits": {0: "batch"}},
opset_version=17, dynamo=False, external_data=False,
)
print("FP32 export complete", flush=True)
def fp16(output: Path) -> None:
import onnx
from onnxconverter_common import float16
graph = onnx.load(output / "onnx/model.onnx")
graph = float16.convert_float_to_float16(graph, keep_io_types=True)
onnx.save(graph, output / "onnx/model_fp16.onnx")
print("FP16 conversion complete (integer inputs and FP32 logits retained)", flush=True)
def int8(output: Path) -> None:
from onnxruntime.quantization import QuantType, quantize_dynamic
quantize_dynamic(
str(output / "onnx/model.onnx"), str(output / "onnx/model_int8.onnx"),
weight_type=QuantType.QInt8, per_channel=True, reduce_range=False,
op_types_to_quantize=["MatMul"],
extra_options={"MatMulConstBOnly": True},
)
print("INT8 dynamic MatMul conversion complete (embeddings retained in FP32)", flush=True)
def reference(source: Path, output: Path) -> None:
import numpy as np
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
torch.set_num_threads(4)
tokenizer = AutoTokenizer.from_pretrained(source, local_files_only=True)
model = AutoModelForSequenceClassification.from_pretrained(
source, local_files_only=True, dtype=torch.float32,
attn_implementation="eager",
).eval()
cases = json.loads((output / "validation/fixtures.json").read_text())
arrays, metadata = {}, []
with torch.inference_mode():
for case in cases:
inputs = tokenizer(case["texts"], padding=True, truncation=True,
max_length=512, return_tensors="pt")
logits = model(**inputs).logits.cpu().numpy()
key = case["id"]
arrays[key + "_input_ids"] = inputs["input_ids"].numpy()
arrays[key + "_attention_mask"] = inputs["attention_mask"].numpy()
arrays[key + "_logits"] = logits
metadata.append({"id": key, "shape": list(inputs["input_ids"].shape)})
print("Reference", key, metadata[-1]["shape"], flush=True)
np.savez_compressed(output / "validation/reference.npz", **arrays)
(output / "validation/reference.json").write_text(json.dumps({
"upstream": UPSTREAM, "revision": REVISION, "precision": "float32",
"attention_implementation": "eager", "provider": "PyTorch CPU",
"cases": metadata, "purpose": "Numerical conversion checks; not a labeled accuracy benchmark.",
"source_weights_sha256": digest(source / "model.safetensors"),
"fixtures_sha256": digest(output / "validation/fixtures.json"),
"reference_npz_sha256": digest(output / "validation/reference.npz"),
}, indent=2) + "\n")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("stage", choices=["export", "fp16", "int8", "reference"])
parser.add_argument("--source", type=Path, required=True)
parser.add_argument("--output", type=Path, default=Path(__file__).resolve().parents[1])
args = parser.parse_args()
(args.output / "onnx").mkdir(parents=True, exist_ok=True)
if args.stage in {"export", "reference"}:
verify_source(args.source, args.output)
if args.stage == "export":
export(args.source, args.output)
elif args.stage == "fp16":
fp16(args.output)
elif args.stage == "int8":
int8(args.output)
else:
reference(args.source, args.output)