File size: 2,763 Bytes
d176ecd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
#!/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()