skip to content

What do prepare_pt2e and convert_pt2e do to an exported PyTorch graph?

level: middleimportance: must knowfreq 60%

answer

  1. three beats: annotate, observe, rewrite
  2. one call inserts, one call rewrites
  3. the data pass is yours to run
  4. observers record ranges, nothing is int8 yet
  5. scale and zero-point are computed at convert

basics

~20 s

prepare_pt2e inserts observer modules at the tensors a quantizer annotated, so a calibration or training run records their value ranges. convert_pt2e then replaces each observer with quantize/dequantize op pairs carrying the scale and zero-point it computed.

solid answer

~50 s

They are the two halves of a graph rewrite, with a data pass in between. `prepare_pt2e(graph, quantizer)` walks the exported graph, and wherever the quantizer annotated a node — say a conv whose weights and output should be int8 — it splices in an **observer**: a small module that watches tensors flowing through and accumulates their min/max or a histogram. You then run the prepared module on representative inputs; nothing is quantized yet, the numbers are still float, the observers are just recording. `convert_pt2e(prepared)` reads each observer, turns the recorded range into a concrete `scale` and `zero_point`, deletes the observer, and inserts `quantize_per_tensor`/`dequantize_per_tensor` (or per-channel) pairs around the op. Weights get folded to int8 at this point. The result is still a runnable graph — it simulates int8 numerics in float — that you export again and lower.

code

python · 22 lines
python
import 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_pt2e, convert_pt2e

model = nn.Sequential(nn.Conv2d(3, 16, 3), nn.ReLU(), nn.Flatten(), nn.Linear(16 * 30 * 30, 10)).eval()
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)
)

prepared = prepare_pt2e(graph, quantizer)   # observers in, still float
for _ in range(32):                          # the beat you must not skip
    prepared(torch.randn(1, 3, 32, 32))

quantized = convert_pt2e(prepared)           # observers out, q/dq in
print(quantized(*example_inputs).shape)

go deeper

for a junior

Remember the order: prepare, then a forward-only pass over sample inputs, then convert. Say plainly that no integers exist until convert_pt2e runs.

for a middle

Explain what an observer is and how a recorded min/max becomes a scale and zero-point, and why weights can be per-channel while activations are usually per-tensor.

for a senior

Demonstrate the diagnostic habits: inspect the converted graph, compare float and quantized intermediate outputs, and treat a missing or unrepresentative calibration pass as a first suspect when accuracy drops.

for a principal

Own calibration as a pipeline asset — a versioned, representative dataset with a quality gate — so that requantizing after every model retrain is routine rather than a hand-run notebook.

## The shape of the flow PT2 export quantization is a **three-beat rhythm**: annotate, observe, rewrite. The quantizer does the annotating, `prepare_pt2e` sets up the observing, your calibration loop performs it, and `convert_pt2e` does the rewriting. Knowing which beat does what is the single most useful thing to have straight, because almost every bug in this pipeline is a beat that did not happen. ## Beat one: prepare_pt2e inserts observers An exported graph is a list of operations with tensors flowing between them. The quantizer has already tagged some of those tensors with a **quantization spec**: target dtype (usually `torch.int8`), symmetric or affine, per-tensor or per-channel, the qmin/qmax range, and which observer class to use. `prepare_pt2e` materialises those tags. For every tagged tensor it inserts an observer module into the graph at that point. An observer is deliberately dumb: it takes a tensor, updates its running statistics, and returns the tensor unchanged. Common ones are `MinMaxObserver` (track the smallest and largest value ever seen), `MovingAverageMinMaxObserver` (an EMA of those, less sensitive to one freak batch) and `HistogramObserver` (build a value histogram and pick a range that minimises quantization error, which tolerates outliers better but costs more time). After `prepare_pt2e`, the model still computes in float32 and produces exactly the same outputs it did before — plus a small runtime overhead from the observers. That is what makes the next step safe. ## Beat two: you supply the data The framework will not calibrate for you. You run the prepared module yourself: ``` for batch in calibration_batches: prepared(batch) ``` No labels, no loss, no `backward()` — you only need the forward pass so observers see realistic activations. (If you used `prepare_qat_pt2e` instead, this beat is a real fine-tuning loop and the inserted modules are fake-quant modules, not plain observers.) Skipping this beat is the classic silent failure: observers that never saw data hold degenerate ranges, and the converted model returns garbage or zeros without raising anything. ## Beat three: convert_pt2e rewrites the graph `convert_pt2e` walks the observed graph and, for each observer, computes the affine mapping from the observed float range to the integer range. For symmetric int8 that is essentially `scale = max(abs(min), abs(max)) / 127` with `zero_point = 0`; for affine schemes the zero-point shifts so that float 0 maps exactly onto an integer. It then deletes the observer and inserts a quantize/dequantize pair — the `quantize_per_tensor` / `dequantize_per_tensor` ops (or their per-channel variants) — around the operation, with the computed scale and zero-point baked in as constants. Weight tensors are quantized eagerly and stored as integers. Activations get q/dq nodes on the wire. The converted graph is still executable in Python and still returns float tensors at its boundaries. What changed is that every quantized value now round-trips through the integer grid, so the numbers you see are the numbers the device will compute. That is why you evaluate accuracy **after** `convert_pt2e`, never on the prepared model: the prepared model is still float and will flatter you. ## What happens next, and why it matters here The q/dq pairs are not the end state. When you export the converted graph and lower it, the backend's partitioner looks for those pairs around patterns it can execute natively and fuses them into a genuine int8 kernel — the q/dq nodes disappear into the kernel's input/output handling. Where the backend cannot claim a pattern, the q/dq pair survives as literal work: quantize, dequantize, then run the op in float. That is strictly slower than not quantizing, and it is why the annotation step is backend-specific rather than generic. ## Failure modes worth naming - **No calibration run.** Observers hold default ranges; accuracy collapses, nothing errors. - **Calibration data unlike production data.** Ranges too narrow, so real inputs clip; or too wide, so resolution is wasted on values that never occur. - **Evaluating the prepared model.** It is float; the numbers mean nothing about the deployed model. - **Per-tensor weights on a conv with wildly different channel magnitudes.** One large channel sets the scale for all of them and the small channels round to zero. `is_per_channel=True` in the config exists exactly for this. - **Re-running prepare on an already-converted graph.** The q/dq nodes are ordinary ops now; the second pass will happily observe them and produce nonsense. ## The compact answer prepare_pt2e = insert observers where the quantizer said to. Your loop = let them see real data. convert_pt2e = turn recorded ranges into scales and zero-points, drop the observers, insert q/dq ops, fold weights to int8. Everything after that is lowering.

  • If you evaluate accuracy on the prepared model instead of the converted one, what do you see?
    Essentially the float model's accuracy. Observers only record; they do not round anything. The prepared model computes in float32 all the way through, so it cannot show you quantization error. Any accuracy check has to run on the output of convert_pt2e, which simulates the integer grid.
  • What actually differs between MinMaxObserver and HistogramObserver?
    MinMaxObserver takes the extreme values it saw, so one outlier batch stretches the range and wastes resolution on values that almost never occur. HistogramObserver builds a distribution and picks a clipping range that minimises overall quantization error, tolerating outliers at the cost of a slower calibration pass.
  • Why does convert_pt2e leave quantize/dequantize nodes in the graph rather than emitting int8 kernels directly?
    Because the framework does not know which kernels the target has. The q/dq pairs are a portable, backend-neutral encoding of the numerics. During lowering the backend partitioner matches the patterns it supports and folds the q/dq nodes into real int8 kernels; anything it cannot claim keeps the explicit pair and runs in float.
  • Can you inspect what the flow decided, before shipping it?
    Yes — print the converted graph. You can see which ops sit between quantize/dequantize pairs, whether the scales are per-tensor or per-channel constants, and where the quantized region starts and stops. A quantized region that ends far earlier than you expected is the usual sign the quantizer skipped a pattern.

saying these in an interview costs you the question

  • Believes prepare_pt2e already produces an int8 model
  • Skips the calibration forward pass entirely
  • Measures accuracy on the prepared rather than converted model
  • Thinks scales are chosen by the framework without data
  • Assumes the q/dq nodes remain as overhead after lowering

context