face-model / code /quantize_int8.py
Banaxi-Tech's picture
Upload face detector (YOLO11n/s), ONNX exports, scripts, model card
d176ecd verified
Raw History Blame Contribute Delete
2.76 kB
#!/usr/bin/env python3
"""Post-training static INT8 quantization (QDQ, per-channel weights) of the face detector ONNX.
Usage: python quantize_int8.py export/face_yolo11n_fp32.onnx export/face_yolo11n_int8.onnx [--calib 300]
"""
import argparse
import random
from pathlib import Path
import cv2
import numpy as np
import onnx
from onnxruntime.quantization import CalibrationDataReader, CalibrationMethod, QuantFormat, QuantType, quantize_static
def letterbox(img, size=640):
h, w = img.shape[:2]
s = min(size / h, size / w)
nh, nw = round(h * s), round(w * s)
img = cv2.resize(img, (nw, nh), interpolation=cv2.INTER_LINEAR)
canvas = np.full((size, size, 3), 114, np.uint8)
top, left = (size - nh) // 2, (size - nw) // 2
canvas[top:top + nh, left:left + nw] = img
return canvas
class Reader(CalibrationDataReader):
def __init__(self, files, input_name):
self.it = iter(files)
self.name = input_name
def get_next(self):
f = next(self.it, None)
if f is None:
return None
img = letterbox(cv2.imread(str(f)))[:, :, ::-1].transpose(2, 0, 1) # BGR->RGB, CHW
return {self.name: (img[None].astype(np.float32) / 255.0)}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("src")
ap.add_argument("dst")
ap.add_argument("--calib", type=int, default=300)
ap.add_argument("--images", default="dataset/images/train")
a = ap.parse_args()
random.seed(0)
files = sorted(Path(a.images).glob("*.jpg"))
random.shuffle(files)
files = files[:a.calib]
g = onnx.load(a.src).graph
inp = g.input[0].name
# The Detect head's decode (DFL box decoding, class Sigmoid, final Concats) mixes box values up to 640 with
# 0-1 class scores in one tensor; INT8 would give them one shared scale and crush the scores. Keep it float.
exclude = [n.name for n in g.node
if n.name.startswith("/model.23/") and not n.name.startswith(("/model.23/cv2.", "/model.23/cv3."))]
print(f"keeping {len(exclude)} decode nodes in float")
quantize_static(a.src, a.dst, Reader(files, inp), quant_format=QuantFormat.QDQ, per_channel=True,
weight_type=QuantType.QInt8, activation_type=QuantType.QUInt8,
calibrate_method=CalibrationMethod.MinMax, nodes_to_exclude=exclude)
# keep ultralytics metadata (class names, stride, imgsz) so the model loads like the original
src, q = onnx.load(a.src), onnx.load(a.dst)
del q.metadata_props[:]
q.metadata_props.extend(src.metadata_props)
onnx.save(q, a.dst)
print(f"RES wrote {a.dst}: {Path(a.dst).stat().st_size / 1e6:.1f} MB (from {Path(a.src).stat().st_size / 1e6:.1f} MB)")
if __name__ == "__main__":
main()