cone-distance / pt_quantization.py
Aryan Sethi
Claude Opus 5 (1M context)
Quantise the ONNX export, and record how the model performs
35206ab
Raw History Blame Contribute Delete
4.71 kB
"""Quantise the PyTorch checkpoint, and measure whether it was worth it.
The short answer is that it is not, and the reason is worth knowing rather than
taking on trust -- so this quantises the model and benchmarks the result instead
of asserting anything.
PyTorch offers two routes on CPU:
* **Dynamic quantisation** is one line and needs no calibration, but it only
converts operators whose weights can be quantised without knowing the
activation range: Linear, LSTM, GRU, Embedding. YOLO11n is convolutions --
81 Conv and 7 ConvolutionDepthWise in the exported graph, and a handful of
Linear layers at most. So it converts almost nothing.
* **Static quantisation** does cover convolutions, but requires the model to be
built for it: QuantStub and DeQuantStub around the graph, fusable Conv-BN-ReLU
patterns, and a forward that traces cleanly. Ultralytics' model is none of
those things, and making it so means forking the architecture.
Which is why the deployment path for this project is export-then-quantise --
ONNX, NCNN, TFLite -- rather than quantising the checkpoint. The exporters
lower the graph to plain convolutions first, and the runtimes have kernels that
actually execute INT8 convolutions quickly. `onnx_quantization.py` is that path,
and it gets 4.5x smaller at no measurable accuracy cost.
python pt_quantization.py --weights trained/yolo11n/weights/best.pt
"""
from __future__ import annotations
import argparse
import time
from pathlib import Path
import numpy as np
import torch
def count_layers(module: torch.nn.Module) -> dict[str, int]:
counts: dict[str, int] = {}
for layer in module.modules():
name = type(layer).__name__
if name in ("Conv2d", "Linear", "BatchNorm2d", "ConvTranspose2d"):
counts[name] = counts.get(name, 0) + 1
return counts
def time_forward(model: torch.nn.Module, tensor: torch.Tensor,
runs: int, warmup: int) -> tuple[float, float]:
"""p50 and p95 milliseconds for a raw forward pass.
Raw forward rather than Ultralytics' predict(), so the measurement is the
model and not the pre/post-processing around it -- those are identical
between the two variants and would dilute the comparison.
"""
with torch.inference_mode():
for _ in range(warmup):
model(tensor)
samples = []
for _ in range(runs):
started = time.perf_counter()
model(tensor)
samples.append((time.perf_counter() - started) * 1000.0)
values = np.array(samples)
return float(np.percentile(values, 50)), float(np.percentile(values, 95))
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--weights", default="trained/yolo11n/weights/best.pt")
parser.add_argument("--imgsz", type=int, default=640)
parser.add_argument("--runs", type=int, default=30)
parser.add_argument("--warmup", type=int, default=5)
parser.add_argument("--threads", type=int, default=4)
args = parser.parse_args()
from ultralytics import YOLO
torch.set_num_threads(args.threads)
model = YOLO(args.weights).model.float().eval()
tensor = torch.zeros(1, 3, args.imgsz, args.imgsz)
print(f"layers: {count_layers(model)}")
quantised = torch.ao.quantization.quantize_dynamic(
model, {torch.nn.Linear, torch.nn.LSTM, torch.nn.GRU}, dtype=torch.qint8)
converted = sum(1 for layer in quantised.modules()
if "quantized" in type(layer).__module__)
print(f"layers actually converted to int8: {converted}")
original = Path(args.weights).stat().st_size
out = Path(args.weights).with_name(Path(args.weights).stem + "_dynint8.pt")
torch.save(quantised.state_dict(), out)
print(f"\nsize fp32 {original / 1e6:5.1f} MB "
f"dynamic int8 {out.stat().st_size / 1e6:5.1f} MB")
fp32 = time_forward(model, tensor, args.runs, args.warmup)
int8 = time_forward(quantised, tensor, args.runs, args.warmup)
print(f"\nforward pass at {args.imgsz}px, {args.threads} threads, "
f"{args.runs} runs")
print(f" fp32 p50 {fp32[0]:7.2f} ms p95 {fp32[1]:7.2f} ms")
print(f" dynamic int8 p50 {int8[0]:7.2f} ms p95 {int8[1]:7.2f} ms")
print(f" speedup {fp32[0] / int8[0]:.2f}x")
if converted == 0:
print("\nNothing was converted. Dynamic quantisation covers Linear, LSTM,")
print("GRU and Embedding; this graph is convolutions. Use the export")
print("path instead -- see onnx_quantization.py.")
if __name__ == "__main__":
main()