Skip to main content

Quantization-Aware Training

Quantization-aware training (QAT) fine-tunes a PyTorch model while simulating the numerical effects of INT8 inference. Use it when post-training quantization causes an unacceptable accuracy loss and you can retrain the model with representative data.

SiMa QAT is designed to run inside the model's existing PyTorch training project. It does not replace the dataset, augmentations, loss function, optimizer, or validation metric. Keeping those parts of the original project is important because they are usually what made the floating-point model accurate in the first place.

Start from a pretrained floating-point checkpoint when possible. Training from random initialization is supported, but usually requires substantially more time and data.

How QAT works​

Preparation adds observers and fake-quantization operations to the model. Observers measure activation ranges, while fake quantization rounds and clamps values during the forward pass to approximate INT8 execution. Tensors, gradients, and optimizer updates remain floating point, allowing training to adapt the model weights to those quantization effects.

The workflow is:

  1. Prepare the eager PyTorch model for QAT.
  2. Warm up observers using the normal training loop.
  3. Freeze the activation ranges and SiMa-compatible weight scales.
  4. Recover accuracy by continuing to train with the locked scales.
  5. Finalize the model for inference.
  6. Export a standard opset-17 ONNX model containing QuantizeLinear and DequantizeLinear (QDQ) nodes.

Install​

The QAT wheel requires Python 3.10 or newer and PyTorch 2.8.x. Install it in the environment that already contains the model's training dependencies.

Download the QAT package with sima-cli:

sima-cli neat install qat

The command downloads the wheel and installs or refreshes the QAT coding-agent skill for Codex and Claude. It does not change the active Python environment. Activate the training environment and install the downloaded wheel:

python -m pip install ./sima_qat-*.whl
python -c "import torch, sima_qat; print(torch.__version__, sima_qat.__file__)"

Add QAT to a training project​

The following steps form one continuous workflow. Adapt the model, data, optimizer, loss, and validation calls to the existing training project.

1. Prepare the model​

Prepare the model before constructing the optimizer. Preparation returns an isolated QAT graph; it does not modify or move the source model or example inputs. The input tuple must match the model's positional inputs, dtypes, and shapes.

import torch
from sima_qat import (
sima_export_onnx,
sima_finalize_qat_model,
sima_freeze_qat,
sima_prepare_qat_model,
)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# Recreate the model and load the floating-point checkpoint.
source_model = MyModel()
source_model.load_state_dict(torch.load("model-fp32.pt", map_location="cpu"))
source_model.train()

example_inputs = (torch.randn(1, 3, 224, 224),)
qat_model = sima_prepare_qat_model(
source_model,
example_inputs,
device=device,
)

optimizer = torch.optim.AdamW(qat_model.parameters(), lr=1e-5)
criterion = torch.nn.CrossEntropyLoss()

Build the optimizer from qat_model, not source_model, because training updates the prepared graph.

2. Train, freeze, and recover​

Train normally at first so the observers can measure representative activation ranges. Freeze the quantization parameters after this warm-up, then continue training so the model can recover accuracy with locked quantization grids.

During validation, keep fake quantization enabled but temporarily disable observers so held-out data cannot change their ranges. eval() and inference_mode() alone do not stop observers. Restore their previous states in finally, including observers already disabled by freezing.

from torch.ao.quantization import disable_observer
from torch.ao.quantization.fake_quantize import FakeQuantizeBase

freeze_epoch = 2
num_epochs = 4

for epoch in range(num_epochs):
qat_model.train()

# Reserve one or more later epochs for recovery training.
if epoch == freeze_epoch:
sima_freeze_qat(qat_model)

for images, labels in train_loader:
images = images.to(device)
labels = labels.to(device)

optimizer.zero_grad(set_to_none=True)
predictions = qat_model(images)
loss = criterion(predictions, labels)
loss.backward()
optimizer.step()

observer_states = [
(module, module.observer_enabled.clone())
for module in qat_model.modules()
if isinstance(module, FakeQuantizeBase)
]
try:
qat_model.apply(disable_observer)
validate(qat_model, validation_loader, device)
finally:
for module, enabled in observer_states:
module.observer_enabled.copy_(enabled)

The freeze epoch is model-dependent. A useful starting point is to warm up observers for most of a short fine-tuning run and reserve at least one final epoch for recovery. Track validation accuracy before and after freezing. If accuracy drops sharply, freeze earlier and allow more recovery training.

3. Save and resume training​

Save the prepared model before finalization so training can be resumed. A checkpoint should contain both the QAT model and optimizer state.

from pathlib import Path

checkpoint_dir = Path("checkpoints")
checkpoint_dir.mkdir(parents=True, exist_ok=True)
torch.save(
{
"epoch": epoch,
"model": qat_model.state_dict(),
"optimizer": optimizer.state_dict(),
},
checkpoint_dir / f"qat-{epoch:02d}.pt",
)

To resume, recreate and prepare the same model with the same example-input and batch contract before loading the saved states:

checkpoint = torch.load("checkpoints/qat-03.pt", map_location="cpu")

source_model = MyModel()
source_model.load_state_dict(torch.load("model-fp32.pt", map_location="cpu"))
qat_model = sima_prepare_qat_model(
source_model,
example_inputs,
device=device,
)
optimizer = torch.optim.AdamW(qat_model.parameters(), lr=1e-5)

qat_model.load_state_dict(checkpoint["model"])
optimizer.load_state_dict(checkpoint["optimizer"])
start_epoch = checkpoint["epoch"] + 1

The checkpoint preserves observer state, the frozen state, and the exact quantization parameters used by finalization and ONNX export.

4. Finalize and export​

Finalization creates an inference-only model. Move the trained QAT model and example inputs to CPU, then export the finalized model as opset-17 QDQ ONNX.

final_model = sima_finalize_qat_model(qat_model.cpu())
export_inputs = tuple(value.cpu() for value in example_inputs)
sima_export_onnx(
final_model,
export_inputs,
"model.qdq.onnx",
input_names=["images"],
output_names=["predictions"],
device="cpu",
)

Batch size​

Preparation keeps the leading tensor dimension dynamic by default. This lets the same prepared graph train with ordinary data-loader batch sizes, handle a short final batch, and export with the concrete batch size passed to sima_export_onnx.

Most models need no batch option. Set dynamic_batch=False only when the model intentionally requires the exact example batch size:

qat_model = sima_prepare_qat_model(
source_model,
example_inputs,
device=device,
dynamic_batch=False,
)

Fixed-batch capture is appropriate when the model asserts or branches on batch size, uses fixed-size recurrent state, or folds batch into direction, channel, or other layout geometry. Dynamic capture compares the captured output with the original model on the supplied example and fails with guidance to disable dynamic batch when it would change the model's behavior.

Validate and compile​

Measure the task-level metric for the original floating-point model, the prepared model before and after freezing, the finalized PyTorch model, and the ONNX model. This makes it clear which lifecycle step introduced a regression.

Validate the exported file before compilation:

python - <<'PY'
import onnx

model = onnx.load("model.qdq.onnx")
onnx.checker.check_model(model)
print("ONNX model is valid")
PY

Compare the finalized PyTorch model with ONNX Runtime on representative validation samples. Small elementwise differences can occur at quantization boundaries, so use tolerances appropriate to the outputs and confirm the model's real accuracy metric.

Common weighted, activation, normalization, pooling, reduction, and shape/layout operations are covered. ArgMax and TopK keep index outputs as integers. PReLU, ConvTranspose, Embedding/Gather, GridSample, ReduceMin, and CumSum remain trainable but are not QAT-annotated in this release.

The exported QDQ ONNX model is the handoff to Model Compiler. Import, partitioning, optimization, and hardware assignment are separate compilation steps.

Runnable examples​

The repository's examples include a small CPU-friendly MNIST workflow, ImageNet fine-tuning of a pretrained classifier, and a plain-PyTorch YOLO26n workflow with checkpoint resume and export.