The opt-in mini_trainer.modeling.quantization API uses TorchAO PT2E for
post-training calibration (PTQ) and quantization-aware training (QAT).
Converted inference executes oneDNN integer Conv/Linear kernels with static INT8
weights and UINT8 activations. QAT uses fake quantization and float32 master
parameters/gradients; it does not promise reduced training memory or integer
backward computation. Native CUDA INT8 training has a
separate implementation and qualification scope.
Install the optional dependency in an explicitly selected backend environment:
uv sync --extra cpu --extra quantization
# Existing CUDA environments: do not run a CPU sync; select their CUDA extra.The recorded x86 qualification uses PyTorch 2.12 and TorchAO 0.17. TorchAO is loaded only when this API is used; ordinary training and prediction are unchanged.
import torch
from mini_trainer.modeling.quantization import prepare_int8, load_int8
# model is a loaded floating-point mini_trainer model. All inputs below are
# float32 CPU batches AFTER the same preprocessing used for ordinary inference.
prepared = prepare_int8(model, example_batch)
with torch.no_grad():
for images in training_calibration_batches:
prepared(images)
converted = prepared.convert()
converted.save(
"int8-model",
example_batch,
preprocessing={"recipe": "record the actual resize, scale and normalization"},
calibration={"split": "train", "manifest_sha256": "record the actual manifest hash"},
)
inference, coverage = load_int8("int8-model").lower(example_batch)
with torch.no_grad():
scores = inference(example_batch)Calibration must use training data, never held-out validation/test examples. The caller supplies provenance; the API cannot infer the provenance of tensors. The prepared PTQ graph is a calibration object, including when its mode is eval. Convert it before evaluating held-out data. Conversion refuses unobserved or nonfinite ranges, and does not modify the prepared model.
Weights use symmetric per-channel int8; activations use affine per-tensor uint8.
The default reduce_range=True uses activation values 0..127 (seven effective
bits in eight-bit storage). This avoids intermediate saturation in oneDNN's
AVX2/non-VNNI kernels. Full range 0..255 is available with reduce_range=False
for a validated VNNI deployment target. The strict reference-versus-native
parity check remains enabled in both modes; tolerances are unchanged.
This portability fix changes the default calibration/QAT range. Recreate older
full-range QAT checkpoints with reduce_range=False to preserve their recipe.
Existing exported full-range graphs are not silently recalibrated on load;
they still require compatible hardware and successful lowering parity.
New reduced-range recipes record the activation bounds explicitly.
Bias, normalization, score transforms and other non-linear operations may remain
floating point. The report includes the actual remaining operator inventory.
All captured Conv1d/Conv2d/Linear operations must receive weight and activation
annotations. Lowering fails if it cannot produce integer kernels or leaves
floating Conv/Linear kernels. It never labels a plain Q/DQ reference execution
as native integer inference.
The bundle contains a reference model.pt2 graph, checksum, input shape, class
metadata, structured output mapping, bit widths, dependency versions,
preprocessing/calibration provenance, and verified lowering coverage. Real calibration tensors are excluded from the saved program. Packing
is performed again on the deployment CPU. A reference graph alone is not an
accelerated runtime. Existing output directories are never overwritten.
prepared = prepare_int8(model, example_batch, qat=True)
optimizer = torch.optim.AdamW(prepared.parameters(), lr=1e-4)
prepared.train()
for images, targets in training_batches:
optimizer.zero_grad()
loss = criterion(prepared(images), targets)
loss.backward()
optimizer.step()
prepared.freeze_observers() # Optional: hold learned ranges fixed for later steps.
prepared.eval() # Evaluation does not update QAT ranges or BatchNorm.
with torch.no_grad():
scores = prepared(validation_batch)
converted = prepared.convert()Construct the optimizer after preparation. Save prepared.state_dict() and
the optimizer/scheduler/scaler states. Restore into an identically prepared model,
then restore optimizer state. Observer ranges, fake-quant flags and the explicit
freeze flag are part of the state. The recipe is checked on restoration. This is
not an ordinary Classifier.build(weights=...) checkpoint: automatic reconstruction
through mt_train/mt_predict remains a subsequent integration step.
QAT runs can use train_one_epoch with an appropriate criterion, disabled EMA,
and no embedding-dependent regularizer. Captured graphs neither consume ambient
supervision nor populate EmbeddingContext. Train/eval switching covers dropout
and BatchNorm; arbitrary Python training branches are specialized by capture.
Autoregressive teacher-forcing/sampling requires a separate training integration
and is not currently a supported QAT claim. Capture failures propagate explicitly.
Inputs currently have a fixed captured batch and image shape. All batches must match it; pad and slice final inference batches, or prepare a separate shape. The original model's parameters, modes and caches are preserved by preparation. Functional linears, weight parametrization, hierarchical aggregation and class masks are included in capture; this does not rely on a backbone allowlist.
EMA remains unsupported. Model quality, memory and latency require measured comparisons beyond compatibility checks. See the quantization roadmap for planned work and benchmark findings for retained measurements and their separate backend/workload boundaries.
Run focused checks without changing the installed environment:
OMP_NUM_THREADS=1 bash dev/check.sh test tests/quantization/test_quantization.pyBackend references: PT2E x86 quantization,
QAT workflow.
Strict graph capture is intentional: the installed backend otherwise misses
functional-linears' source metadata. Explicit lower_pt2e_quantized_to_x86
provides native kernels without relying on a compiler silently optimizing Q/DQ.