Quantization-aware training for Cortex-M0+
A practical walkthrough of moving quantization noise from inference into training, and what it took to get an 8 KB keyword-spotting model below 4 KB without losing accuracy on a $4 MCU.
The goal was simple: a yes — detect — stop keyword spotter that runs on an STM32G0 with 32 KB of flash and 8 KB of RAM, and fits inside the budget of a consumer device that already costs a few dollars. The naive approach — train in float32, post-training quantize to int8 — left us at 9 KB. We needed 4.
Why post-training quantization stopped being enough
PTQ is wonderful when it works. It is fast, requires no retraining, and adds zero complexity to your training pipeline. For models that have headroom, it is the right tool.
Our keyword spotter did not have headroom. The architecture was a small CNN with two conv layers and a dense head, and the dense weights were already sparse-ish from a previous round of pruning. PTQ threw away the small but real signal sitting in the long tail of the weight distribution.
Quantization noise is most painful exactly where the model has been most cleverly pruned.
The QAT recipe, simplified
The version of QAT we ended up using is, in essence, three steps:
- Train the model in float32 as usual, to a baseline accuracy.
- Add fake-quantization nodes to the forward pass: weights and activations are quantized to int8 and immediately dequantized back to float32. Gradients still flow through, but the model learns to be robust to quantization noise.
- Continue training with the fake-quant nodes in place. The loss surface changes — you are now optimizing for accuracy under quantization.
The training code, stripped down
We used PyTorch's built-in torch.quantization utilities for the fake-quant
nodes, but the rest of the loop is plain PyTorch. The interesting bit is the calibration
step — we used a small held-out set of yes / no / stop recordings from our own
hardware to compute the activation ranges.
import torch
from torch.quantization import prepare_qat, convert
# 1. Start with a float32 model
model = SmallKWS().to(device)
# 2. Fuse conv+bn+relu where the runtime supports it
model.fuse_model()
# 3. Insert fake-quant nodes
model.qconfig = torch.quantization.get_default_qat_qconfig("qnnpack")
prepare_qat(model, inplace=True)
# 4. Train normally — gradient flows through fake-quant
for epoch in range(EPOCHS):
train_one_epoch(model, train_loader, opt)
eval_model(model, val_loader)
# 5. Convert to a real int8 model
model.eval()
int8_model = convert(model, inplace=False)
# 6. Hand to tinymlr for on-device deployment
int8_model.export("kws_int8.tiny")
What surprised us
The first surprise was that the size win was modest — about 1 KB. Most of the model was already small. The bigger surprise was the accuracy win. PTQ top-1 on our test set was 91.2%. QAT jumped to 93.8%. Two and a half points for free on a model that we had already pruned.
The second surprise was where the gains came from. We expected them in the dense layer (the one most prone to quantization error). They came from the first conv layer — the one that processes raw audio. Apparently, learning to ignore the rounding noise on a 16-bit PCM input makes the entire downstream stack more robust.
The on-device side
On the device, the int8 model runs in 1.4 ms on the Cortex-M0+ at 64 MHz. Flash footprint: 3.6 KB. RAM at peak: 2.1 KB. Both inside budget, with room for the rest of the firmware.
The full pipeline from .py to .bin is now repeatable end-to-end
in under four minutes, including quantization. If you want to try it, the
tinymlr repo has the export step and a small example.
Next week: why your model isn't the slow part. I'll be profiling a similar audio inference loop on STM32H7 and showing where the time actually goes.