Quantization Aware Training (QAT)

Overview

Quantization Aware Training (QAT) is a technique for improving the accuracy of models which are quantized by applying “fake” quantizations to the model’s weights (and optionally, activations) during training. This fake quantization allows for the model to adjust for noise introduced by the quantization, so when the model is eventually quantized, the accuracy loss is minimized. We use the quantization techniques implemented in torchao to provide support for QAT and post-training quantization (PTQ) in axolotl.

We recommend reviewing the excellent QAT tutorial in the torchtune library, and the QAT documentation in the torchao library, for more details.

Configuring QAT in Axolotl

To enable QAT in axolotl, add the following to your configuration file:

qat:
  activation_dtype: # Optional[str] = "int8". Fake quantization layout to use for activation quantization. Valid options are "int4", "int8", "float8"
  weight_dtype: # Optional[str] = "int8". Fake quantization layout to use for weight quantization. Valid options are "int4", "fp8", "nvfp4", and "ternary".
  group_size: # Optional[int] = 32. The number of elements in each group for per-group fake quantization
  fake_quant_after_n_steps: # Optional[int] = None. The number of steps to apply fake quantization after

We support the following quantization schemas:

  • Int4WeightOnly (requires the fbgemm-gpu extra when installing Axolotl)
  • Int8DynamicActivationInt4Weight
  • Float8DynamicActivationFloat8Weight
  • Float8DynamicActivationInt4Weight
  • NVFP4
  • Ternary (BitNet b1.58 style)

Ternary QAT

weight_dtype: ternary trains every linear outside the LM head with weights restricted to {-1, 0, 1}, scaled by a per-row (per output channel) absmean, as in BitNet b1.58. Norms and the LM head stay in high precision, and quantize_embedding uses int8 per-row rather than ternary. group_size does not apply.

Activation quantization is opt-in: set activation_dtype: int8 to also fake quantize activations with a per-token absmax scale (int8 is the only activation dtype ternary accepts). Leave activation_dtype unset — the default — for weight-only ternary.

The use_onebitllms path trains the same scheme through an external library, for checkpoints that are already ternary.

qat:
  weight_dtype: ternary
  activation_dtype: int8        # optional, omit for weight-only; int8 is the only value ternary accepts
  quantize_embedding: true      # int8 per-row embeddings
  fake_quant_after_n_steps: 100 # optional high precision warmup

Gradients reach the high precision master weights through a straight-through estimator, so the optimizer state stays unquantized. At the end of training the ternary values are baked into the saved weights, which means the checkpoint loads with plain transformers but is still stored at the model’s training dtype — serving it at two bits per weight requires an inference stack with a ternary format (for example GGUF TQ1_0/TQ2_0). Recovering quality after switching an existing model to ternary takes a continued pretraining scale token budget, not a fine-tuning one.

Once you have finished training, you must quantize your model by using the same quantization configuration which you used to train the model with. You can use the quantize command to do this.