Int8 PTQ cost you 4 points of accuracy on device — how do you recover it?
answer
- climb a ladder, do not jump
- measure the converted model, on device
- most damage sits in a few layers
- one flag fixes per-tensor weight crush
- the last rung needs labels and a training loop
basics
~20 sWork cheapest-first: verify calibration coverage, switch weights to per-channel, then find the sensitive layers by comparing float and quantized intermediates and exempt them. Only if that falls short move to QAT with prepare_qat_pt2e and a short fine-tune.
solid answer
~50 sTreat it as a debugging ladder, not a jump to QAT. **First**, confirm you are measuring the converted model on device, not the prepared float one, and that the calibration set really matches production inputs — narrow ranges clip and that alone can cost several points. **Second**, set `is_per_channel=True`; per-tensor weight scales let one large channel crush the small ones. **Third**, localise the damage: run float and quantized versions side by side and compare intermediate tensors, then exempt the worst offenders with `set_module_name`, trading a little speed for accuracy. **Only then** escalate to QAT: build the quantizer with `is_qat=True`, call `prepare_qat_pt2e`, fine-tune at a low learning rate for a small fraction of the original schedule so the weights adapt to rounding, then `convert_pt2e` and re-measure. QAT costs a training pipeline, labels and revalidation on every retrain.
code
python · 25 linesimport torch
from torch import nn
from executorch.backends.xnnpack.quantizer.xnnpack_quantizer import (
XNNPACKQuantizer,
get_symmetric_quantization_config,
)
from torchao.quantization.pt2e.quantize_pt2e import prepare_qat_pt2e, convert_pt2e
model = nn.Sequential(nn.Conv2d(3, 8, 3), nn.ReLU(), nn.Flatten(), nn.Linear(8 * 30 * 30, 10))
example_inputs = (torch.randn(1, 3, 32, 32),)
graph = torch.export.export(model, example_inputs).module()
quantizer = XNNPACKQuantizer().set_global(
get_symmetric_quantization_config(is_per_channel=True, is_qat=True)
)
qat_model = prepare_qat_pt2e(graph, quantizer)
opt = torch.optim.SGD(qat_model.parameters(), lr=1e-4)
for _ in range(10): # short fine-tune at a low LR
loss = qat_model(torch.randn(1, 3, 32, 32)).sum()
loss.backward()
opt.step()
opt.zero_grad()
quantized = convert_pt2e(qat_model)go deeper
Know that accuracy loss after quantization is normal and has cheap first fixes — better calibration data and per-channel weights — before anything involving retraining.
Explain what QAT changes: fake-quant modules apply rounding in the forward pass so the loss sees it, while gradients pass through, letting weights adapt to the integer grid.
Show the diagnostic sequence and the discipline behind it: measure the converted model on device, localise damage per layer, exempt selectively, and only then pay for QAT.
Set the policy: an accuracy floor and latency budget agreed before work starts, and a clear-eyed view of whether the team can carry QAT as a permanent pipeline stage across every retrain.
## Do not start with QAT QAT is the expensive answer: it needs the training data, a working fine-tuning loop, GPU time, and a commitment to repeat all of it every time the float model is retrained. A four-point drop is frequently caused by something upstream that costs an afternoon to fix. Work the ladder in order. ## Rung 1: check that you measured the right thing Two mistakes produce fake numbers. Evaluating the **prepared** model tells you nothing — it is still float32, observers only record. And evaluating the converted model **on the host** can disagree with the device, because a delegate's kernels are not bit-identical to the simulated numerics. Establish the real baseline: converted model, on the target device, on your real eval set. ## Rung 2: audit the calibration set Static quantization is only as good as the ranges it observed. Ask what production inputs look like versus what you calibrated on — resolution, preprocessing, lighting, speakers, class balance. Ranges that are too narrow clip real inputs at qmin/qmax; ranges stretched by a single freak batch waste the integer grid. Two fixes: widen the calibration set to cover the real distribution, and use a histogram-style observer instead of plain min/max so a lone outlier stops dictating the scale. This rung alone often recovers most of the loss. ## Rung 3: per-channel weights With per-tensor weight quantization, one scale covers an entire convolution kernel. If channel 7 has weights ten times larger than channel 3, the shared scale is set by channel 7 and channel 3's weights round to a handful of distinct values — sometimes to zero. Per-channel gives each output channel its own scale, costs almost nothing at inference on kernels that support it, and is why `get_symmetric_quantization_config(is_per_channel=True)` is the sane default. If you shipped per-tensor, flip this before doing anything harder. ## Rung 4: localise, then exempt Quantization damage is rarely uniform; usually two or three layers cause most of it. Run the float and converted models on the same inputs and compare intermediate activations layer by layer — a signal-to-quantization-noise ratio per layer, or just relative error. Typical offenders are the first layer (raw input dynamic range), the final classifier or logits head, and any block with heavy outliers. Having found them, exempt them via `set_module_name(name, ...)` on the quantizer — a different config, or none at all — and re-run prepare/convert. You give back some speed and some size in exchange for accuracy, and because you know *which* layers you exempted, that trade is explicit rather than accidental. ## Rung 5: QAT If the ladder runs out and you still miss the accuracy floor, QAT is the tool that genuinely changes the answer, because it lets the weights adapt to rounding instead of being rounded after the fact. Mechanically it mirrors PTQ. Build the quantizer with `get_symmetric_quantization_config(is_per_channel=True, is_qat=True)`, call `prepare_qat_pt2e` instead of `prepare_pt2e`, and instead of a forward-only calibration loop you run a real fine-tuning loop with labels and an optimizer. The inserted **fake-quant** modules round the tensor in the forward pass — so the loss reflects the integer grid — while letting gradients pass through, so training can still make progress. A few epochs at a low learning rate, typically a small fraction of the original schedule, is normal; QAT is fine-tuning, not retraining. Then `convert_pt2e` produces the real quantized graph. Two cautions. **Evaluate the converted model, not the QAT model** — the fake-quant graph is still float and can look better than what ships. And **QAT is not a one-off**: every subsequent float retrain has to go back through the QAT loop, so you are committing to it as a permanent pipeline stage, with the data access, compute budget and revalidation cost that implies. ## What a good answer sounds like Name the ladder and the reasoning behind its order: measurement first because it is free, calibration second because it is the most common cause, per-channel third because it is a one-flag fix, selective exemption fourth because it converts an accuracy problem into a controlled speed trade, and QAT last because it is the only rung with an ongoing organisational cost. Mentioning that you would set an accuracy floor and a latency budget *before* starting — so you know when to stop climbing — is what separates a senior answer from a list of techniques.
- How do you find which layers are responsible for the accuracy loss?Run the float and converted models on identical inputs and compare intermediate activations layer by layer, using a signal-to-quantization-noise ratio or relative error. The distribution is usually skewed: a couple of layers dominate. Suspect the input layer, the logits head, and any block whose activations carry heavy outliers.
- Why does fake quantization let gradients flow if rounding has zero derivative almost everywhere?Because the backward pass uses a straight-through estimator: rounding is applied in the forward pass so the loss sees the integer grid, while the gradient is passed through as if the rounding were the identity (clipped outside the representable range). Training can then adjust weights to be robust to the rounding it experiences.
- What is the ongoing cost of committing to QAT?It becomes a permanent pipeline stage. Every float retrain must re-run the QAT fine-tune, which needs training data access, labels, compute, and a fresh on-device accuracy revalidation. Teams underestimate this; a PTQ pipeline that recovers accuracy through better calibration is far cheaper to keep green.
- Why can host-side evaluation of the converted model disagree with on-device results?The converted graph simulates integer numerics using the reference implementation, while the device runs the delegate's own kernels — different accumulation order, fusion, and occasionally different rounding. The gap is usually small, but the number that matters is the one measured on the target hardware with the shipped artifact.
saying these in an interview costs you the question
- Jumps straight to QAT before checking calibration
- Reports accuracy from the prepared, still-float model
- Leaves weights per-tensor and blames quantization
- Retrains from scratch instead of a short fine-tune
- Validates the fake-quant model rather than the converted one