NewYour coding agent can read the release notes before it upgrades.Set up the MCP server →
PyPI · #2748 most downloaded on PyPI
Package for applying ao techniques to GPU models
Last release 2 months ago
03 Aug 2026
Ships fairly regularly
a new release about every 2 months
Most releases are documented
notes for 16 of 22 stable releases
Nothing withdrawn
no release was ever pulled
3 years old
22 releases · first in 2023
The deprecated v1 AffineQuantizedTensor stack and layout system have been fully removed, completing the migration to v2 tensor subclasses.
NVFP4 for dense linears is now supported for training (prototype), implementing the recipe from Pretraining Large Language Models with NVFP4. This is a prototype feature and the API may change without notice. Hardening, MoE support, and torchtitan integration are planned.
Example usage:
import torch
from torchao.quantization import quantize_
from torchao.prototype.moe_training.nvfp4_training.nvfp4_training import (
NVFP4TrainingConfig,
)
# Blackwell (SM100+); bfloat16 model with Linear dims divisible by 128
model = model.cuda().bfloat16()
# Replace nn.Linear with NVFP4Linear
quantize_(model, NVFP4TrainingConfig())
# Train as usual
out = model(x)
out.sum().backward()
Support for PyTorch older than 2.11 has been dropped from the code, and CI has moved to newer PyTorch versions.
The deprecated v1 AffineQuantizedTensor stack and layout system have been fully removed, completing the migration to v2 tensor subclasses.
check_cpu_version and check_xpu_version helpers by @Xia-Weiwen in https://github.com/pytorch/ao/pull/4211cc1 typo in Float8LinearConfig.__post_init__ operand-precision assertion by @GotFusion in https://github.com/pytorch/ao/pull/4115== typo in float8_transpose axiswise dim remapping by @GotFusion in https://github.com/pytorch/ao/pull/4038Int4WeightOnlyConfig.group_size in fake quant configs by @oriollinan in https://github.com/pytorch/ao/pull/4518nn.Linear from #4587 by @lisjin in https://github.com/pytorch/ao/pull/4613dim from kwargs in Float8Tensor aten.split handler by @GotFusion in https://github.com/pytorch/ao/pull/4429quant_lift_up to CustomGraphPass with uuid() by @frgossen in https://github.com/pytorch/ao/pull/4424RopeSDPAFusionPass(CustomGraphPass) with uuid() by @frgossen in https://github.com/pytorch/ao/pull/4423rowwise_scaled_linear_sparse_cutlass_f8f8.cu by @rgommers in https://github.com/pytorch/ao/pull/4426Full Changelog: https://github.com/pytorch/ao/compare/v0.17.0...v0.18.0-rc1
One column per month.
Deprecate AQT and related classes
We are excited to announce the 0.17 release of torchao! This release adds support for cuteDSL MXFP8 MoE kernels, per-head FP8 quantized low precision attention, ABI stability, and more!
We added a new CuteDSL MXFP8 quantization kernel for 3d expert weights that writes scale factors directly to blocked layout for tensorcores: https://github.com/pytorch/ao/pull/4090
We added a new API for per-head fp8 quantized attention with FA3 as the backend (https://github.com/pytorch/ao/pull/3959 and https://github.com/pytorch/ao/pull/3857)
Example Usage of Direct Replacement:
from torchao.prototype.attention.fp8_fa3 import fp8_fa3_sdpa, fp8_fa3_rope_sdpa
out = fp8_fa3_sdpa(q, k, v)
Example Usage of Wrapper:
from torchao.prototype.attention import (
AttentionBackend,
LowPrecisionAttentionConfig,
apply_low_precision_attention,
)
# Instantiate any nn.Module()
model = MyModel()
# Simple SDPA replacement
config = LowPrecisionAttentionConfig(backend=AttentionBackend.FP8_FA3)
model = apply_low_precision_attention(model, config)
# Flash activation is handled internally by the wrapper
output = model(inputs)
# Torch.compile will enable rope fusion
model = torch.compile(model)
As of https://github.com/pytorch/ao/issues/3516, torchao is now ABI stable for all cuda kernels! This means if the user is running torch 2.11+, they will be able to access torchao’s cuda kernels without having to upgrade torch version for each new torchao version. This applies to the current and all future torchao releases (0.17.0+). Note that python-only API compatibility is the same as before: we support the latest 3 torch minor versions.
Before:
# Compatible versions for cpp extensions:
torchao 0.16.0 + torch 2.10.0
torchao 0.15.0 + torch 2.9.1
torchao 0.14.1 + torch 2.9.0
# Compatible versions for python-only API
# 3 most recent torch versions:
torchao 0.16.0 + torch 2.10.0, 2.9.1, 2.8.0
torchao 0.15.0 + torch 2.9.1, 2.8.0, 2.7.1
torchao 0.14.0 + torch 2.9.0, 2.8.0, 2.7.0
After:
# Compatible versions for cpp extensions (all future releases):
torchao 0.17.0 onwards + any torch 2.11+
# Compatible versions for python-only API
# 3 most recent torch versions (same as before)
torchao 0.17.0 + torch 2.11.0, 2.10.0, 2.9.1
We are planning to remove these classes in the future, so we are deprecating them in advance. Here are all the classes that now have new deprecation warnings:
Base classes (torchao/dtypes/utils.py):
1. Layout
2. PlainLayout
3. AQTTensorImpl
torchao/dtypes/affine_quantized_tensor.py:
4. AffineQuantizedTensor
torchao/dtypes/floatx/:
5. Float8Layout
6. Float8AQTTensorImpl
7. CutlassSemiSparseLayout
8. CutlassSemiSparseTensorImpl
torchao/dtypes/uintx/:
9. TensorCoreTiledLayout
10. TensorCoreTiledAQTTensorImpl
11. SemiSparseLayout
12. SemiSparseAQTTensorImpl
13. Int4CPULayout
14. Int4CPUAQTTensorImpl
15. Int4XPULayout
16. Int4XPUAQTTensorImpl
17. QDQLayout
18. QDQTensorImpl
19. PlainAQTTensorImpl
20. PackedLinearInt8DynamicActivationIntxWeightLayout
21. PackedLinearInt8DynamicActivationIntxWeightAQTTensorImpl
This was added for torch 2.5 as a placeholder for torch.int4. Now all torch versions we support have these dtypes, so this class is no longer needed and should not be used by anyone. We will remove it next release (0.18.0).
Before:
TorchAODType.INT1
TorchAODType.INT2
TorchAODType.INT3
TorchAODType.INT4
TorchAODType.INT5
TorchAODType.INT6
TorchAODType.INT7
After:
torch.int1
torch.int2
torch.int3
torch.int4
torch.int5
torch.int6
torch.int7
GemliteUIntXWeightOnlyConfig is now replaced with torchao.prototype.UIntxWeightOnlyConfig and Int8DynamicActivationUIntxWeightConfig
# 0.16:
GemliteUIntXWeightOnlyConfig(group_size, bit_width, packing_bitwidth, mode="weight_only")
GemliteUIntXWeightOnlyConfig(group_size, bit_width, packing_bitwidth, mode="dynamic")
# 0.17
UIntxWeightOnlyConfig(group_size, bit_width, packing_bitwidth)
Int8DynamicActivationUIntxWeightConfig(group_size, bit_width, packing_bitwidth)
choose_qparams_affine_with_min_max (https://github.com/pytorch/ao/pull/3808)_emulated_mxfp8_scaled_grouped_mm_2d_2d torch.compile compatible (https://github.com/pytorch/ao/pull/3906)torch._higher_order_ops.scan in pt2e quantization (https://github.com/pytorch/ao/pull/3882)torch._higher_order_ops.while_loop in pt2e quantization (https://github.com/pytorch/ao/pull/3916)Full Changelog: https://github.com/pytorch/ao/compare/v0.16.0...v0.17.0-rc1
…Blocks for Training with Expert Parallelism and deprecated older versions of some configs and less used quantization options to keep torchao leaner! W…
We are excited to announce the 0.16.0 release of torchao! This release adds support for MXFP8 MoE Building Blocks for Training with Expert Parallelism and deprecated older versions of some configs and less used quantization options to keep torchao leaner! We also revamped our doc page, README and made some progress in making torchao ABI stable.
This release includes the following differentiable building blocks for MXFP8 MoE Training with expert parallelism:
These autograd functions can be chained together to implement efficient MoE training with expert parallel comms and grouped GEMMs in MXFP8.
This approach achieves 10% - 25% tokens/second speedup for DeepSeekV3 16b training:
Float8WeightOnlyConfig, Float8DynamicActivationFloat8WeightConfig, Int8DynamicActivationIntxWeightConfig, IntxWeightOnlyConfig, Int4WeightOnlyConfig (https://github.com/pytorch/ao/pull/3510, https://github.com/pytorch/ao/pull/3511, https://github.com/pytorch/ao/pull/3512, https://github.com/pytorch/ao/pull/3513)# v0.15.0 - version 1 was available.
config = Float8WeightOnlyConfig(version=1, ...)
config = Float8DynamicActivationFloat8WeightConfig(version=1, ...)
config = Int8DynamicActivationIntxWeightConfig(version=1, ...)
config = IntxWeightOnlyConfig(version=1, ...)
config = Int4WeightOnlyConfig(version=1, ...)
# v0.16.0 - use version 2 (default). Using version 1 is no longer supported.
config = Float8WeightOnlyConfig(version=2, ...)
config = Float8DynamicActivationFloat8WeightConfig(version=2, ...)
config = Int8DynamicActivationIntxWeightConfig(version=2, ...)
config = IntxWeightOnlyConfig(version=2, ...)
config = Int4WeightOnlyConfig(version=2, ...)
Int8DynamicActivationInt4WeightConfig, Int4DynamicActivationInt4WeightConfig, GemliteUIntXWeightOnlyConfig, Float8StaticActivationFloat8WeightConfig, UIntXWeightOnlyConfig, FPXWeightOnlyConfig to prototype (https://github.com/pytorch/ao/pull/3491)# v0.15.0
from torchao.quantization import (
Int8DynamicActivationInt4WeightConfig,
Int4DynamicActivationInt4WeightConfig,
GemliteUIntXWeightOnlyConfig,
Float8StaticActivationFloat8WeightConfig,
UIntXWeightOnlyConfig,
# removed in this release
FPXWeightOnlyConfig,
)
# v0.16.0. These workflows may be deleted in a future version
from torchao.prototype.quantization.quant_api import (
Int8DynamicActivationInt4WeightConfig,
GemliteUIntXWeightOnlyConfig,
Float8StaticActivationFloat8WeightConfig,
UIntXWeightOnlyConfig,
# removed in this release
FPXWeightOnlyConfig,
Int4DynamicActivationInt4WeightConfig,
)
Int4DynamicActivationInt4WeightConfig, FPXWeightOnlyConfig, Float8DynamicActivationFloat8SemiSparseWeightConfig and SRELUFloat8SemiSparseDynamicActivationFLoat8WeightConfig (https://github.com/pytorch/ao/pull/3723, https://github.com/pytorch/ao/pull/3520, https://github.com/pytorch/ao/pull/3744)# v0.15.0
config = Int4DynamicActivationInt4WeightConfig()
config = FPXWeightOnlyConfig(3, 2)
config = fpx_weight_only(3, 2) # deprecated alias
config = Float8DynamicActivationFloat8SemiSparseWeightConfig()
config = SRELUFloat8SemiSparseDynamicActivationFloat8WeightConfig()
quantize_(model, config)
# v0.16.0
#
The configs are dropped. Please use torchao <= 0.15.0 to use these configs.
CutlassInt4PackedLayout, MarlinQQQLayout, CutlassSemiSparseLayout (https://github.com/pytorch/ao/pull/3723, https://github.com/pytorch/ao/pull/3612, https://github.com/pytorch/ao/pull/3613, https://github.com/pytorch/ao/pull/3744)# v0.15.0
## CutlassInt4PackedLayout
config = Int8DynamicActivationInt4WeightConfig(
group_size=None,
mapping_type=MappingType.SYMMETRIC,
act_mapping_type=MappingType.SYMMETRIC,
layout=CutlassInt4PackedLayout(),
)
## MarlinQQQLayout
config = Int8DynamicActivationInt4WeightConfig(layout=MarlinQQQLayout())
config = Int4DynamicActivationInt4WeightConfig(layout=MarlinQQQLayout())
## MarlinSparseLayout
apply_fake_sparsity(model)
config = Int4WeightOnlyConfig(layout=MarlinSparseLayout(), version=1)
config = Int4WeightOnlyConfig(int4_packing_format="marlin_sparse", version=2)
## CutlassSemiSparseLayout
config = Float8DynamicActivationFloat8SemiSparseWeightConfig(layout=CutlassSemiSparseLayout())
## quantizing the model with the config
quantize_(model, config)
# v0.16.0
# The config and layout options are dropped. Please use torchao <= 0.15.0 to use the layout
# v0.15.0
from torchao.quantization import MultiTensorInputRecorder, Int4WeightOnlyGPTQQuantizer
model = get_model()
input_recorder = MultiTensorInputRecorder()
for i in range(calibration_limit):
args = get_next_input()
input_recorder(*args)
quantizer = Int4WeightOnlyGPTQQuantizer()
args = input_recorder.get_recorded_inputs()
quantizer.quantize(model, *args)
args = get_next_input()
out = model(*args)
# v0.16.0 - this functionality is deleted. We are starting over
# for GPTQ in https://github.com/pytorch/ao/tree/main/torchao/prototype/gptq
# v0.15.0 from torchao.quantization.smoothquant import (
swap_linear_with_smooth_fq_linear,
smooth_fq_linear_to_inference,
)
swap_linear_with_smooth_fq_linear(model_copy, alpha=0.75)
model_copy(**encoded_input)
smooth_fq_linear_to_inference(model_copy)
# v0.16.0 - the functionality above is deleted. Use `torchao.prototype.smoothquant`
# instead (https://github.com/pytorch/ao/tree/main/torchao/prototype/smoothquant)
torchao.autoquant (https://github.com/pytorch/ao/pull/3741)This functionality will be removed in a future release of torchao. Please see https://github.com/pytorch/ao/issues/3739 for context.
wgrad_with_hp option for mxfp8 moe training (https://github.com/pytorch/ao/pull/3508)a2a_dispatch autograd function (https://github.com/pytorch/ao/pull/3579)permute autograd func + triton kernels (https://github.com/pytorch/ao/pull/3580)unpermute autograd function (https://github.com/pytorch/ao/pull/3581)wgrad_with_hp recipe to quantize_ api (https://github.com/pytorch/ao/pull/3611)emulated mode (https://github.com/pytorch/ao/pull/3724)_to_mxfp8_then_scaled_grouped_mm wrapper that accepts keyword args (https://github.com/pytorch/ao/pull/3561)scaled_grouped_mm support gfx942 fp8 data type (https://github.com/pytorch/ao/pull/3540)torch._dynamo.nonstrict_trace if it exists in the torch version (https://github.com/pytorch/ao/pull/3650)Float8SemiSparseTensor off of AQT (https://github.com/pytorch/ao/pull/3361)permute op for FP8 conv (https://github.com/pytorch/ao/pull/3533)rowwise_scaled_linear_sparse_cutlass ABI stable (https://github.com/pytorch/ao/pull/3725)HQQ option in UIntxWeightOnlyConfig (https://github.com/pytorch/ao/pull/3829)RCEIL in triton_to_mxfp8_dim0 kernel with inline PTX for mxfp8 (https://github.com/pytorch/ao/pull/3498)Int8Tensor (https://github.com/pytorch/ao/pull/3489)Float8Tensor (https://github.com/pytorch/ao/pull/3526)scaled_embedding_bag for CPU (https://github.com/pytorch/ao/pull/3755)FqnToConfig handle module swap configs (https://github.com/pytorch/ao/pull/3492)F.scaled_mm (https://github.com/pytorch/ao/pull/3675)Float8DynamicActivationFloat8SemiSparseWeightConfig (https://github.com/pytorch/ao/pull/3595)FqnToConfig to API reference quantization page (https://github.com/pytorch/ao/pull/3709)get_current_accelerator_device (https://github.com/pytorch/ao/pull/3634)Full Changelog: https://github.com/pytorch/ao/compare/v0.15.0...v0.16.0-rc1
Add deprecation warnings for various inference configs
We are excited to announce the 0.15.0 release of torchao! This release adds:
Training runs on 64 node GB200 cluster with TorchTitan Llama4 Scout demonstrated a 1.2x e2e training speedup with equivalent convergence to bfloat16 training baseline. In fact, after 3,000 steps it finishes with slightly lower loss than bfloat16! This is consistent with our scaling experiments with MXFP8 training for dense models.
<img width="790" height="403" alt="mxfp8_with_loss" src="https://github.com/user-attachments/assets/42d54b96-caa6-4ff7-8bac-16541cd8d7f8" />
| Number of GPUs | BF16 tokens/sec | MXFP8 tokens/sec | MXFP8 speedup vs BF16 |
|---|---|---|---|
| 512 | 6169 | 7401 | 1.20x |
See the TorchAO MXFP8 MoE training documentation for more details. You can also check out the TorchTitan MXFP8 documentation to run pretraining jobs with TorchAO MXFP8 by adding a single config.
You can now save and load TorchAO model checkpoints using safetensors! This feature is integrated with Hugging Face transformers starting from v5.0.0 and vLLM 0.13.0 for model inference/serving.
We currently support the following stable configs:
Float8DynamicActivationFloat8WeightConfig
Int4WeightOnlyConfig
IntxWeightOnlyConfig
Int8DynamicActivationIntxWeightConfig
Int8WeightOnlyConfig
Int8DynamicActivationInt8WeightConfig
and will continue to add support for configs as they become stable in the future.
Example:
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, TorchAoConfig
from torchao.quantization import Float8WeightOnlyConfig
model_id = "facebook/opt-125m"
quant_config = Float8WeightOnlyConfig()
quantization_config = TorchAoConfig(quant_type=quant_config)
quantized_model = AutoModelForCausalLM.from_pretrained(
model_id,
device_map="auto",
torch_dtype=torch.bfloat16,
quantization_config=quantization_config,
)
tokenizer = AutoTokenizer.from_pretrained(model_id)
#### Push to hub
MODEL_NAME = model_id.split("/")[-1]
save_to = f"torchao-testing/{MODEL_NAME}-Float8WeightOnlyConfig-v2-0.15.0.dev-safetensors"
quantized_model.push_to_hub(save_to, safe_serialization=True)
tokenizer.push_to_hub(save_to)
To serve a safetensors model checkpoint on vLLM, use the CLI command vllm serve torchao-testing/Qwen3-8B-INT4-0.15.0dev-safetensors.
Now we support running int4 weight only with preshuffled packing format in vllm. We can ship int4 weight only quantized checkpoint with plain packing format (e.g. https://huggingface.co/torchao-testing/opt-125m-Int4WeightOnlyConfig-v2-0.14.0.dev), and the Tensor with plain format will be converted to preshuffled format automatically in SM90+ and when fbgemm_gpu_genai is available.
FqnToConfig (#3083)import torch
from torchao.quantization import quantize_, Float8WeightOnlyConfig, FqnToConfig
class Example(torch.nn.Module):
def __init__(self):
super().__init__()
self.custom_name = torch.nn.Parameter(torch.Tensor([1]))
model = Example()
config = FqnToConfig({"custom_name": Float8WeightOnlyConfig()})
# filter_fn is disabled when quantizing by fqn
quantize_(model, config, filter_fn=None)
print(model.custom_name)
#Float8Tensor(self.act_quant_kwargs=None, self.qdata=tensor([448.], dtype=torch.float8_e4m3fn), self.scale=tensor([0.0022]), self.block_size=[1], self.mm_config=None, self.kernel_preference=<KernelPreference.AUTO: 'auto'> self.shape=torch.Size([1]), self.device=device(type='cpu'), self.dtype=torch.float32)
This is to better support MoE models (Llama4, Deepseek, gpt-oss) which store their weights under attributes like "down_proj", "gate_up_proj", etc. Previously we assumed that the weights were always stored under "weights" as we were primarily focused on quantizing nn.Linear layers.
Currently parameter quantization is only enabled for Int4WeightOnlyConfig, Float8DynamicActivationFloat8WeightConfig, and Float8WeightConfig.
int4_weight_only (https://github.com/pytorch/ao/pull/3145)Before:
from torchao.quantization import (
float8_dynamic_activation_float8_weight,
float8_static_activation_float8_weight,
float8_weight_only,
fpx_weight_only,
gemlite_uintx_weight_only,
int4_dynamic_activation_int4_weight,
int4_weight_only,
int8_dynamic_activation_int4_weight,
int8_dynamic_activation_int8_weight,
int8_weight_only,
quantize_,
uintx_weight_only,
)
quantize_(model, float8_dynamic_activation_float8_weight())
quantize_(model, float8_static_activation_float8_weight(torch.randn(3)))
quantize_(model, float8_weight_only())
quantize_(model, fpx_weight_only(3, 2))
quantize_(model, gemlite_uintx_weight_only())
quantize_(model, int4_dynamic_activation_int4_weight())
quantize_(model, int4_weight_only())
quantize_(model, int8_dynamic_activation_int4_weight())
quantize_(model, int8_dynamic_activation_int8_weight())
quantize_(model, int8_weight_only())
quantize_(model, uintx_weight_only(torch.uint4))
After:
from torchao.quantization import (
Float8DynamicActivationFloat8WeightConfig,
Float8StaticActivationFloat8WeightConfig,
Float8WeightOnlyConfig,
FPXWeightOnlyConfig,
GemliteUIntXWeightOnlyConfig,
Int4DynamicActivationInt4WeightConfig,
Int4WeightOnlyConfig,
Int8DynamicActivationInt4WeightConfig,
Int8DynamicActivationInt8WeightConfig,
Int8WeightOnlyConfig,
quantize_,
UIntXWeightOnlyConfig,
)
quantize_(model, Float8DynamicActivationFloat8WeightConfig())
quantize_(model, Float8StaticActivationFloat8WeightConfig(torch.randn(3)))
quantize_(model, Float8WeightOnlyConfig())
quantize_(model, FPXWeightOnlyConfig(3, 2))
quantize_(model, GemliteUIntXWeightOnlyConfig())
quantize_(model, Int4DynamicActivationInt4WeightConfig())
quantize_(model, Int4WeightOnlyConfig())
quantize_(model, Int8DynamicActivationInt4WeightConfig())
quantize_(model, Int8DynamicActivationInt8WeightConfig())
quantize_(model, Int8WeightOnlyConfig())
quantize_(model, UIntXWeightOnlyConfig(torch.uint4))
Before:
model = torch.nn.Sequential(
torch.nn.Linear(128, 128),
torch.nn.Linear(128, 128),
torch.nn.Conv2d(128, 128, 3, 1, 1),
).cuda().to(torch.bfloat16)
config = ModuleFqnToConfig({
"0": Float8DynamicActivationFloat8WeightConfig(),
})
# these are equivalent
quantize_(model, config, filter_fn=_is_linear)
quantize_(model, config, filter_fn=None)
quantize_(model, config)
After:
# VALID: user must specify None
quantize_(model, config, filter_fn=None)
# INVALID: these now error!
quantize_(model, config, filter_fn=_is_linear)
quantize_(model, config)
Before:
from torchao.utils import TORCH_VERSION_AT_LEAST_2_8
if TORCH_VERSION_AT_LEAST_2_8:
print("PyTorch version was 2.8+")
After: (no replacement)
PerRow(axis) to support axes other than -1 (https://github.com/pytorch/ao/pull/3303)torch.chunk to float8tensor (https://github.com/pytorch/ao/pull/3334)triton_op for dense model (https://github.com/pytorch/ao/pull/3402)extra_repr for nn.Modules (https://github.com/pytorch/ao/pull/3328)int4_weight_only (#3145)" (https://github.com/pytorch/ao/pull/3192)fqn_matches_fqn_config public and fix module device loading (https://github.com/pytorch/ao/pull/3302)enable_fusion_modeling for conv2d and conv3d (https://github.com/pytorch/ao/pull/3343)Full Changelog: https://github.com/pytorch/ao/compare/v0.14.0-rc3...v0.15.0-rc2
…current_default_version=2, please check the deprecation warning warnings.warn( /data/users/jerryzh/ao/torchao/dtypes/uintx/tensor_core_tiled_layout.py…
We are excited to announce the 0.14.1 release of torchao! This release adds support for MoE training on Backwell GPUs and NVFP4 QAT!
We’ve added a quantized building block for speeding up MoE training on Blackwell GPUs: torchao’s `_scaled_grouped_mm`! It is a differentiable drop-in replacement for `torch._grouped_mm` that dynamically quantizes inputs using the given recipe, performs a scaled grouped GEMM, then returns the results in original precision. This results in significant speedups (see benchmarks below)!
import torch
from torch.nn import functional as F
from torchao.prototype.moe_training import (
_scaled_grouped_mm as torchao_scaled_grouped_mm
)
from torchao.prototype.moe_training.conversion_utils import MoEScalingType
from torchao.prototype.moe_training.utils import generate_jagged_offs
num_groups, total_M, N, K = 8, 131072, 8192, 5120
# A = input actvations, B = expert weights
A = torch.randn(total_M, K, dtype=torch.bfloat16, device="cuda", requires_grad=True)
B = torch.randn(num_groups, N, K, dtype=torch.bfloat16, device="cuda", requires_grad=True)
# Token group offsets computed by router in actual MoE layer
offs = generate_jagged_offs(num_groups, total_M, device="cuda")
# Forward and backward example
out = torchao_scaled_grouped_mm(
A,
B.transpose(-2, -1),
offs=offs,
scaling_type=MoEScalingType.MXFP8,
)
labels = torch.ones_like(out)
loss = F.mse_loss(out, labels)
loss.backward()
Microbenchmarks (see README for commands to reproduce benchmarks):
It’s also already integrated into TorchTitan for E2E training with DeepSeekV3 and Llama4! Just use the command line flag: `--model.converters=”quantize.grouped_mm.mx”, which will convert all `torch._grouped_mm` ops to torchao _scaled_grouped_mm ops under the hood:
Torchtitan e2e training benchmarks (see README for commands to reproduce benchmarks):
We added quantization-aware training (QAT) support for NVFP4 as a prototype feature! This feature is currently only available on blackwell GPUs:
from torchao.quantization import quantize_
from torchao.prototype.mx_formats import NVFP4InferenceConfig
# NVFP4 activation and weight quantization
base_config = NVFP4InferenceConfig()
quantize_(model, QATConfig(base_config, step="prepare"))
train(model)
quantize_(model, QATConfig(base_config, step="convert"))
Our NVFP4 QAT support is also integrated into axolotl:
axolotl train examples/llama-3/3b-qat-fsdp2-nvfp4.yaml
Initial experimentation demonstrated that QAT recovered up to 41% mmlu_pro accuracy degradation and 33% bbh accuracy degradation from NVFP4 quantization, for Qwen3-8B and Gemma3-12B-it respectively:
# Qwen3-8B
bbh: 0.779 (baseline) -> 0.7254 (quant) -> 0.7368 (qat), recovered 21.269%
mmlu_pro: 0.4969 (baseline) -> 0.4521 (quant) -> 0.4707 (qat), recovered 41.518%
# Gemma3-12B-it accuracy
bbh: 0.7527 (baseline) -> 0.7068 (quant) -> 0.7222 (qat), recovered 33.551%
mmlu_pro: 0.4074 (baseline) -> 0.3621 (quant) -> 0.3702 (qat), recovered 17.881%
from torchao.quantization import quantize_
from torchao.quantization.qat import (
QATConfig,
Int4WeightPreshuffledFakeQuantizeConfig,
)
# Before
fq_config = Int4WeightPreshuffledFakeQuantizeConfig()
quantize_(m, QATConfig(weight_config=fq_config)
# After
fq_config = Int4WeightFakeQuantizeConfig()
quantize_(m, QATConfig(weight_config=fq_config)
Int4WeightOnlyConfig version from 1 to 2 (https://github.com/pytorch/ao/pull/2949)We updated the implementation for int4 Tensor, so bumps the default version from 1 to 2 for these two configs.
from transformers import AutoModelForCausalLM, AutoTokenizer
model_name = "torchao-testing/opt-125m-Int4WeightOnlyConfig-v1-0.14.dev"
quantized_model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype="bfloat16",
device_map="cuda",
)
/data/users/jerryzh/ao/torchao/core/config.py:250: UserWarning: Stored version is not the same as current default version of the config: stored_version=1, current_default_version=2, please check the deprecation warning
warnings.warn(
/data/users/jerryzh/ao/torchao/dtypes/uintx/tensor_core_tiled_layout.py:241: UserWarning: Models quantized with version 1 of Int4WeightOnlyConfig is deprecated and will no longer be supported in a future release, please upgrade torchao and quantize again, or download a newer torchao checkpoint, see https://github.com/pytorch/ao/issues/2948 for more details
warnings.warn(
Suggestion: upgrade torchao to 0.14 and later and generate the checkpoint again:
quantize_(model, Int4WeightOnlyConfig(group_size=128))
Or download the checkpoint again (please let us know if the checkpoint is not updated)
Please see #2948 for more details around the deprecation.
int4_weight_only (https://github.com/pytorch/ao/pull/2994)The int4_weight_only functions have been superseded by AOBaseConfig objects like Int4WeightOnlyConfig(...), which have been in-use since several previous releases
from torchao.quantization import (
Int4WeightOnlyConfig,
int4_weight_only,
quantize_,
)
# Before
quantize_(m, int4_weight_only())
# After
quantize_(m, Int4WeightOnlyConfig())
# Full list of deprecated functions
float8_dynamic_activation_float8_weight
float8_static_activation_float8_weight
float8_weight_only
fpx_weight_only
gemlite_uintx_weight_only
int4_dynamic_activation_int4_weight
int4_weight_only
int8_dynamic_activation_int4_weight
int8_dynamic_activation_int8_weight
int8_weight_only
uintx_weight_only
Full Changelog: https://github.com/pytorch/ao/compare/v0.13.0...v0.14.0-rc1
Nothing published for this version
Nothing published for this version
\[docs\] Replace deprecated configs with Config objects
We are excited to announce the 0.12.0 release of torchao! This release adds support for QAT + Axolotl Integration and prototype MXFP/NVFP support on Blackwell GPUs!
TorchAO’s QAT support has been integrated into Axolotl’s fine-tuning recipes! Check out the docs here or run it yourself using the following command:
axolotl train examples/llama-3/3b-qat-fsdp2.yaml
axolotl quantize examples/llama-3/3b-qat-fsdp2.yaml
Initial results for Llama3.2-3B by @SalmanMohammadi (https://github.com/axolotl-ai-cloud/axolotl/pull/2590):
| Model/Metric | hellaswag acc | hellaswag acc_norm | wikitext bits_per_byte | wikitext byte_perplexity | wikitext word_perplexity |
|---|---|---|---|---|---|
| bfloat16 | 0.5552 | 0.7315 | 0.6410 | 1.5594 | 10.7591 |
| bfloat16 PTQ | 0.5393 | 0.7157 | 0.6613 | 1.5815 | 11.6033 |
| qat ptq | 0.5423 | 0.7180 | 0.6567 | 1.5764 | 11.4043 |
| Recovered (qat ptq) | 18.87% | 14.56% | 22.66% | 23.08% | 23.57% |
TorchAO now includes prototype support for NVFP4 (NVIDIA's 4-bit floating-point format) and Microscaling (MX) formats on NVIDIA's latest Blackwell GPU architecture. These formats enable efficient inference, achieving up to 61% end-to-end performance improvement in vLLM on Qwen3 models and near 2x speedups for diffusion workloads.
To use:
from torchao.quantization import quantize_
from torchao.prototype.mx_formats import (
MXFPInferenceConfig,
NVFP4InferenceConfig,
)
# Quantize model with MXFP8
model = quantize_(model, MXFPInferenceConfig(block_size=32))
# Quantize model to NVFP4 (without double scaling)
model = quantize_(model, NVFP4InferenceConfig())
Note: This is a prototype feature with APIs subject to change. Requires NVIDIA Blackwell GPUs (B200, 5090) with CUDA 12.8+.
sparsity/prototype/blocksparse (https://github.com/pytorch/ao/pull/2205)AOPerModuleConfig (https://github.com/pytorch/ao/pull/2186)to for fbgemm Tensor (https://github.com/pytorch/ao/pull/2337)Full Changelog: https://github.com/pytorch/ao/compare/v0.11.0...v0.12.0-rc2
We added pytorch 2 export quantization from pytorch to torchao. As part of the planned migration. We’ll follow up with adding deprecation warnings to…
We are excited to announce the 0.11.0 release of torchao! This release adds support for mixture-of-experts (MoE) quantization, PyTorch 2 Export Quantization (PT2E), and a microbenchmarking framework for inference APIs!
We’ve a prototype feature for quantizing MoE modules with a number of TorchAO quantization techniques. This approach leverages the existing TorchAO features for quantizing linear ops and allows them to be used to quantize MoE modules.
from torchao.quantization.prototype.moe_quant.utils import cond_ffn_filter, MoEQuantConfig
from torchao.quantization.quant_api import quantize_, Int8WeightOnlyConfig
quantize_(
model,
MoEQuantConfig(Int8WeightOnlyConfig()),
filter_fn=cond_ffn_filter
)
model=torch.compile(
model,
mode="reduce-overhead",
fullgraph=is_single_token_inference
)
While the above API is all that is needed to quantize a moe module if your moe module is written to be both quantizable and compilable, in practice its rare for a user model to satisfy these conditions due to the variety of MoE implementations. An initial swap of the normal MoE module with a MoEFeedForwardAOQuantizable module is needed to first prepare the model for quantization. An example of this can be found in llama4_quant.py where this technique is demonstrated for the huggingface llama-4-Scout-17B-16E-Instruct model.
We implemented MoE quantization with 2 methods. The first method (designated `base` in the below benchmarks) simply enhances the existing quantized tensor subclass to quantize the 3D MoE expert tensors and perform the necessary indexing and slicing ops while the second method (`fake`), uses a new tensor subclass to simulate a 3D quantized parameter by storing a sequence of 2D slices of the quantized parameter. The first approach is faster with marginally worse memory characteristics. In both cases doing MoE quantization in this way isn’t expected to be maximally performant compared to implementing fused MoE kernels for each technique, but this approach can yield both moderate speedups and significant memory savings.
The following benchmarks are for mixtral-moe run on a single H100 GPU:
| batchsize 1 | batchsize 8 | ||||
|---|---|---|---|---|---|
| Technique | tok/s | memory (GB) | tok/s | tok/s* batch | memory (GB) |
| None | 78.35 | 93.76 | 18.2 | 145.64 | 94.12 |
| int8wo-base | 98.4 | 48.87 | 4.94 | 39.56 | 49.2 |
| int4wo-base | 79.38 | 36.15 | 10.29 | 82.29 | 36.12 |
| fp8wo-base | 59.41 | 52.07 | 2.98 | 23.81 | 52.05 |
| fp8dq-base | 45.92 | 53.97 | 3.78 | 30.23 | 53.94 |
| int8wo-fake | 6.14 | 49.13 | 5.01 | 40.09 | 49.23 |
| int4wo-fake | 14.25 | 30.21 | 11.84 | 94.75 | 30.19 |
| fp8wo-fake | 3.2 | 50.31 | 2.88 | 23.08 | 50.29 |
| fp8dq-fake | 9.78 | 50.92 | 4.08 | 32.61 | 50.89 |
We added pytorch 2 export quantization from pytorch to torchao. As part of the planned migration. We’ll follow up with adding deprecation warnings to PyTorch torch.ao.quantization APIs and updating docs in the future. We also simplified the import path for some of the util functions. Here is a non-exhaustive list of APIs you can use:
# top level APIs
from torchao.quantization.pt2e.quantize_pt2e import prepare_pt2e, prepare_qat_pt2e, convert_pt2e
from torchao.quantization.pt2e.quantizer import X86InductorQuantizer
# export utils
from torchao.quantization.pt2e import (
move_exported_model_to_eval,
move_exported_model_to_train,
allow_exported_model_train_eval
)
# graph utils
from torchao.quantization.pt2e import (
find_sequential_partitions,
get_equivalent_types,
update_equivalent_types_dict,
bfs_trace_with_node_process,
)
# pt2e numeric debugger
from torchao.quantization.pt2e import (
generate_numeric_debug_handle,
CUSTOM_KEY,
NUMERIC_DEBUG_HANDLE_KEY,
prepare_for_propagation_comparison,
extract_results_from_loggers,
compare_results,
)
We’ve introduced a streamlined microbenchmark framework, to help developers track and evaluate the performance of their post-training quantization and sparsity APIs for different matrix sizes and model types. The framework also includes support for advanced GPU and memory profiling techniques, providing deeper insights into performance characteristics.
To run the benchmarks, use the following command:
python -m benchmarks.microbenchmarks.benchmark_runner --config benchmarks/microbenchmarks/test/benchmark_config.yml
Sample Benchmark Results (on 1xH100):
| Name | Quantization | Shape | Baseline Inference Time (ms) | Inference Time (ms) | Speedup |
|---|---|---|---|---|---|
| small_bf16_linear | float8dq-tensor | 16384, 16384, 16384 | 13.34 | 7.72 | 1.73x |
| small_bf16_linear | float8dq-tensor | 16384, 16384, 32768 | 26.04 | 14.62 | 1.78x |
| small_bf16_linear | float8dq-tensor | 16384, 16384, 65536 | 53.59 | 29.05 | 1.84x |
| small_bf16_linear | float8dq-tensor | 16384, 32768, 32768 | 68.94 | 28.07 | 2.46x |
| small_bf16_linear | float8dq-tensor | 16384, 32768, 65536 | 108.63 | 58.7 | 1.85x |
| small_bf16_linear | float8dq-tensor | 16384, 65536, 65536 | 215.66 | 118.42 | 1.82x |
| small_bf16_linear | float8dq-tensor | 32768, 32768, 32768 | 108.16 | 57.09 | 1.89x |
| small_bf16_linear | float8dq-tensor | 32768, 32768, 65536 | 214.74 | 110.08 | 1.95x |
| small_bf16_linear | float8dq-tensor | 32768, 65536, 65536 | 432.44 | 223.46 | 1.94x |
| small_bf16_linear | float8dq-tensor | 65536, 65536, 65536 | 870.37 | 447.97 | 1.94x |
torchao.quantization (https://github.com/pytorch/ao/pull/2134)Full Changelog: https://github.com/pytorch/ao/compare/v0.10.0...v0.11.0
The following usage of \Float8Config\ is deprecated in torchao v0.10.0:
We are excited to announce the 0.10.0 release of torchao! This release adds support for end to end training for mxfp8 on Nvidia B200, PARQ (for quantization aware training), module swap quantization API to for research, and some updates for low bit kernels!
Low bit optimizers (added in 0.4) is moved out of prototype and now have official support in torchao.
We have an early version of the end to end training workflow for the mxfp8 dtypes with torch.compile on NVIDIA B200, with the cuBLAS mxfp8 gemm seeing an observed speedup of over 2x over bfloat16 gemm, and casts from bfloat16 to mxfp8 achieving up to 5.5 TB/s. Please see our README.md for MX for more information. We plan to improve performance further in future releases.
from torchao.prototype.parq.optim import QuantOptimizer, ProxHardQuant
from torchao.prototype.parq.quant import UnifQuantizer
# Separate quantizable from non-quantizable parameter groups
param_groups = [
{"params": weights, "quant_bits": 2}, # add extra quant_bits key for QAT
{"params": others},
]
# Initialize any torch.optim.Optimizer
base_optimizer = torch.optim.SGD(param_groups, lr=0.1, momentum=0.9, weight_decay=1e-4)
# Apply a simple wrapper to quantize in optimizer.step()
optimizer = QuantOptimizer(
base_optimizer, quantizer=UnifQuantizer(), prox_map=ProxHardQuant()
)
We added a prototype API for post-training quantization. Users can swap their linear or embedding layers into their QuantizedLinear and QuantizedEmbedding counterparts, and set the quantizers that specify how they want the input activations or weights to be quantized:
quantized_linear = QuantizedLinear(...)
quantized_linear.weight_quantization = IntQuantizer(
num_bits=4,
group_size=32,
dynamic=True,
quantization_mode="symmetric",
)
quantized_linear.input_quantization = CodeBookQuantizer(
num_bits=8,
features=10,
)
Note: The API is highly subject to change and will be integrated with quantize_ in the future. For more detail, please see the README.
Low-bit CPU and MPS kernels are now pip installable from source. To install torchao with low-bit CPU kernels, you can use the following command on an Arm-based Mac:
USE_CPP=1 pip install git+https://github.com/pytorch/ao.git
You can then quantize your model to run on Arm-based Macs with high-performance CPU kernels in torchao. SharedEmbeddingQuantizer,EmbeddingQuantizer, and Int8DynamicActivationIntxWeightConfig all support 1-8 bit quantization.
from torchao.experimental.quant_api import Int8DynamicActivationIntxWeightConfig, SharedEmbeddingQuantizer, EmbeddingQuantizer
from torchao.quantization.granularity import PerGroup, PerRow
from torchao.quantization.quant_api import quantize_
# Quantize embedding/unembedding to 8-bits with SharedEmbeddingQuantizer
# SharedEmbeddingQuantizer is for quantizing models like Llama1B/3B
# where the embedding/unembedding layers share weights
# If the embedding/unembedding layers do not share weights, use
# EmbeddingQuantizer instead
SharedEmbeddingQuantizer(
weight_dtype=torch.int8,
granularity=PerRow(),
has_weight_zeros=True
).quantize(model) # Quantize linear layers to 4-bits
quantize_(
model,
Int8DynamicActivationIntxWeightConfig(
weight_dtype=torch.int4,
granularity=PerGroup(128),
has_weight_zeros=False,
)
)
The following usage of `Float8Config` is deprecated in torchao v0.10.0:
config = Float8LinearConfig(
cast_config_input=CastConfig(scaling_type=ScalingType.DELAYED),
cast_config_weight=CastConfig(scaling_type=ScalingType.DELAYED),
cast_config_grad_output=CastConfig(scaling_type=ScalingType.DELAYED),
)
If you would like to use float8 training with delayed scaling, please use an earlier release of torchao. Please see https://github.com/pytorch/ao/issues/1680 for more context about this deprecation.
quantize_'s config argument (https://github.com/pytorch/ao/pull/1861)This was done following a deprecation window to simplify the arguments of quantize_, please see https://github.com/pytorch/ao/issues/1690 for more context.
# torchao v.0.9.0
def quantize_(
model: torch.nn.Module,
**config: Union[AOBaseConfig, Callable[[torch.nn.Module], torch.nn.Module]],**
filter_fn: Optional[Callable[[torch.nn.Module, str], bool]] = None,
set_inductor_config: Optional[bool] = None,
device: Optional[torch.types.Device] = None,
):
# torchao v.0.10.0
def quantize_(
model: torch.nn.Module,
config: AOBaseConfig,
filter_fn: Optional[Callable[[torch.nn.Module, str], bool]] = None,
set_inductor_config: Optional[bool] = None,
device: Optional[torch.types.Device] = None,
):
set_inductor_config argument of quantize_. (https://github.com/pytorch/ao/pull/1865)This was done following a deprecation window to decouple quantize_ from torchinductor, please see https://github.com/pytorch/ao/issues/1715 for more context.
# torchao v.0.9.0
def quantize_(
...,
set_inductor_config: Optional[bool] = None,
...,
):
# if set_inductor_config != None, throw a deprecation warning
# if set_inductor_config == None, set it to True to stay consistent with old behavior
# torchao v0.10.0
def quantize_(
...,
):
# set_inductor_config is removed from quantize_ and moved to relevant individual workflows
We removed some of our prototype features that are not used, including DORA (https://github.com/pytorch/ao/pull/1815), split_k kernel (https://github.com/pytorch/ao/pull/1816), profiler (https://github.com/pytorch/ao/pull/1862) and bitnet (https://github.com/pytorch/ao/pull/1866).
QAT
Low Bit Optimizers
Module swap quantization API
Benchmarking
Kernels
AOConfigs
sparsify_ to configs (https://github.com/pytorch/ao/pull/1856)SAM2
QAT
MX
Affine Quantization
Full Changelog: https://github.com/pytorch/ao/compare/v0.9.0...v0.10.0-rc1
…to config: AOBaseConfig, following a deprecation process detailed below.
We are excited to announce the 0.9.0 release of torchao! This release moves a number of sparsity techniques out of prototype, a significant overhaul of the quantize_ api, a new cutlass kernel for 4 bit dynamic quantization and more!
We’ve promoted block sparsity out of torchao.prototype and made several performance improvements. You can accelerate your models with block sparsity as follows:
from torchao.sparsity import sparsify, block_sparse_weight
sparsify_(model, block_sparse_weight(blocksize=64))
| Technique | Decode (tok/s) | Model Size (GB) |
|---|---|---|
| baseline | 134.40 | 15.01 |
| 2:4 sparse | 163.13 | 10.08 |
| bsr-0.8-32 | 210.91 | 6.01 |
| bsr-0.8-64 | 222.43 | 6.00 |
| bsr-0.9-32 | 255.19 | 4.88 |
| bsr-0.9-64 | 262.94 | 4.88 |
| 2:4 sparse + int4wo (marlin) | 255.21 | 3.89 |
Block Sparsity technique names (bsr) indicate sparsity fraction and blocksize.
These numbers were generated on H100 using torchao/_models/llama/generate.py on the Meta-Llama-3.1-8B model. You can reproduce these numbers using this script
W've identified that the binaries are broken on M1 and have been since v0.8.0 though they were working in v0.7.0. We're working on a fix for this, details and discussion can be found here.
We are migrating the way quantize_ workflows are configured from callables (tensor subclass inserters) to direct configuration (config objects). Motivation: align with the rest of the ecosystem, enable inspection of configs after instantiation, remove a common source of confusion.
What is changing:
Specifically, here is how the signature of quantize_'s second argument will change:
#
# torchao v0.8.0 and before
#
def quantize(
model: torch.nn.Module,
apply_tensor_subclass: Callable[[torch.nn.Module], torch.nn.Module],
...,
): ...
#
# torchao v0.9.0
#
def quantize(
model: torch.nn.Module,
config: Union[AOBaseConfig, Callable[[torch.nn.Module], torch.nn.Module]],
...,
): ...
#
# torchao v0.10.0 or later (exact version TBD)
#
def quantize(
model: torch.nn.Module,
config: AOBaseConfig,
...,
): ...
quantize_ changed from apply_tensor_subclass to config. Since the vast majority of callsites today are passing in configuration with a positional argument, this change should not affect most people.quantize_ will change from Callable[[torch.nn.Module], torch.nn.Module] to config: AOBaseConfig, following a deprecation process detailed below.int8_weight_only) to camel case (Int8WeightOnlyConfig). All argument names for each config are kept as-is. We will keep the old snake case names (int8_weight_only) around and alias them to the new names (int8_weight_only = Int8WeightOnlyConfig), to avoid breaking callsites. We plan to keep the old names forever. Here are all the workflow config name changes:| old name (will keep working) | new name (recommended) |
|---|---|
int4_weight_only |
Int4WeightOnlyConfig |
float8_dynamic_activation_float8_weight |
Float8DynamicActivationFloat8WeightConfig |
float8_static_activation_float8_weight |
Float8StaticActivationFloat8WeightConfig |
float8_weight_only |
Float8WeightOnlyConfig |
fpx_weight_only |
FPXWeightOnlyConfig |
gemlite_uintx_weight_only |
GemliteUIntXWeightOnlyConfig |
int4_dynamic_activation_int4_weight |
Int4DynamicActivationInt4WeightConfig |
int8_dynamic_activation_int4_weight |
Int8DynamicActivationInt4WeightConfig |
int8_dynamic_activation_int8_semi_sparse_weight |
n/a (deprecated) |
int8_dynamic_activation_int8_weight |
Int8DynamicActivationInt8WeightConfig |
int8_weight_only |
Int8WeightOnlyConfig |
uintx_weight_only |
UIntXWeightOnlyConfig |
Configuration for prototype workflows using quantize_ will be migrated at a later time.
How these changes can affect you:
quantize_ API workflows and are passing in config by a positional argument (quantize_(model, int8_weight_only(group_size=128))), you are not affected. This positional syntax will keep working going forward. You are encouraged to migrate your callsite to the new config name (quantize_(model, Int8WeightOnlyConfig(group_size=128)) though the old names will continue to work indefinitely.quantize_ API workflows and are passing in config by a keyword argument (quantize_(model, tensor_subclass_inserter=int8_weight_only(group_size=128))), your callsite will break. You will need to change your callsite to quantize_(model, config=int8_weight_only(group_size=128)). We don't expect many people to be in this bucket.quantize_ API, you will need to use the new configuration system. Please see https://github.com/pytorch/ao/issues/1690 for details.sparsify_, you are not affected for now and a similar change will happen in a future version of torchao.This migration will be a two step process:
Please see https://github.com/pytorch/ao/issues/1690 for more details.
Before:
from torchao.prototype.sparsity.superblock.blocksparse import block_sparse_weight
After:
from torchao.sparsity import block_sparse_weight
set_inductor_config argument of quantize_ (https://github.com/pytorch/ao/pull/1716)We are migrating the set_inductor_config argument of quantize_ to individual workflows. Motivation:
quantize_.quantize_ API level, with individual workflows opting in as needed.set_inductor_config to quantize_, your callsite will keep working with a deprecation warning. We recommend that you migrate this option to your individual workflow.set_inductor_config argument will be removed from quantize_.# torchao v0.8.x
def quantize_(
...,
set_inductor_config: bool = True,
...,
): ...
# torchao v.0.9.0
def quantize_(
...,
set_inductor_config: Optional[bool] = None,
...,
):
# if set_inductor_config != None, throw a deprecation warning
# if set_inductor_config == None, set it to True to stay consistent with old behavior
# torchao v TBD (a future release)
def quantize_(
...,
):
# set_inductor_config is removed from quantize_ and moved to relevant individual workflows
Please see https://github.com/pytorch/ao/issues/1715 for more details.
We plan to deprecate delayed and static scaling from torchao.float8 training codebase due to lack of real world use cases for delayed/static scaling (dynamic scaling is required for higher accuracy) and complexity tax for supporting these features.
Supermask (https://pytorch.org/blog/speeding-up-vits/) is a technique for improving the accuracy of block sparsified models by learning a block-sparse mask during a training phase.
from torchao.sparsity import SupermaskLinear, block_sparse_weight
sparsify_(model, lambda x: SupermaskLinear.from_linear(x, block_size=64, sparsity_level=0.9)
# training here
# collapse supermask into a normal linear layer (with many weights set to 0) and then convert to block sparse format for inference speedup
sparsify_(model, lambda x: SupermaskLinear.to_linear(x, sparsity_level=0.9)
sparsify_(model, block_sparse_weight(blocksize=64))
This kernel which adds support for 4 bit dynamic activation + 4 bit weight quantization can be used as follows:
from torchao.quantization import int4_dynamic_activation_int4_weight
quantize_(model, int4_dynamic_activation_int4_weight)
In torchao v0.9.0, we include very early support for training and inference on the NVIDIA Blackwell GPUs following the microscaling recipes from https://www.opencompute.org/documents/ocp-microscaling-formats-mx-v1-0-spec-final-pdf, and backed by real MX gemms.
Here is how to use the current prototype APIs.
:warning: <em> Note that torch.compile support is not fully there yet, there are no guarantees on performance at this time, and we expect to change these APIs rapidly as we iterate in future versions of torchao. Please see https://github.com/pytorch/ao/issues/556 for more details. </em>
MX training
from torchao.prototype.mx_formats.mx_linear import swap_linear_with_mx_linear
from torchao.prototype.mx_formats.config import MXLinearConfig, MXGemmKernelChoice
from torchao.utils import is_sm_at_least_100
# early prototype: on MX-enabled hardware, you can use the real MX gemm backed by
# torchao's CUTLASS kernels. In the future, we will also add cuBLAS kernel support.
gemm_kernel_choice = MXGemmKernelChoice.EMULATED
if is_sm_at_least_100():
gemm_kernel_choice = MXGemmKernelChoice.CUTLASS
m = torch.nn.Sequential(torch.nn.Linear(32, 32)).cuda()
config = MXLinearConfig(
elem_dtype=torch.float8_e4m3fn,
block_size=32,
gemm_kernel_choice=gemm_kernel_choice,
)
swap_linear_with_mx_linear(m, config=config)
# training loop (not shown)
MX inference, weights are in MX and matmul is in high precision.
from torchao.prototype.mx_formats.mx_linear import swap_linear_with_mx_inference_linear
from torchao.prototype.mx_formats.config import MXLinearConfig
m = torch.nn.Sequential(torch.nn.Linear(32, 32)).cuda()
config = MXLinearConfig(elem_dtype=torch.float8_e4m3fn, block_size=32)
swap_linear_with_mx_inference_linear(m, config=config)
# do inference (not shown)
The additional features for MX support in v0.9.0 were enabled by:
mx_mm function and MXLinear. (https://github.com/pytorch/ao/pull/1667).item() issue in running parallel evaluation for BO mixed precision (https://github.com/pytorch/ao/pull/1630)DDP with nf4 (https://github.com/pytorch/ao/pull/1684)ZeroPointDomain.NONE & None zero point domains (https://github.com/pytorch/ao/pull/1556)Full Changelog: https://github.com/pytorch/ao/compare/v0.8.0...v0.9.0-rc1
We are excited to announce the 0.8.0 release of torchao\! In this release we’ve shipped the first CUTLASS kernel in torchAO which adds support for W4A
We are excited to announce the 0.8.0 release of torchao! In this release we’ve shipped the first CUTLASS kernel in torchAO which adds support for W4A8 linear operator. In addition to this, we’ve also added TTFT benchmarks to torchAO and compared different quantization + sparsity speedups for prefill / decoding.
A new W4A8 linear operator is implemented, that corresponds to int8_dynamic_activation_int4_weight quantization where two 4-bit weights get packed into a single 8-bit integer value; also, CUTLASS is made a sub-module of torchao repo, in order to be able to utilize more of its functionality to implement new kernels.
-q parameter |
Average tokens/sec | Average Bandwidth in GB/s | Peak Memory Usage in GB | Model Size in GB |
|---|---|---|---|---|
| 95.24 | 258.55 | 13.90 | 13.21 | |
-q int8wo |
155.31 | 1028.37 | 8.97 | 6.62 |
-q int4wo-32 |
186.70 | 774.98 | 5.31 | 4.15 |
-q int4wo-hqq |
186.47 | 774.01 | 5.04 | 4.15 |
-q int8dq |
49.64 | 328.72 | 9.44 | 6.62 |
-q w4a8-cutlass (tuned) |
119.31 | 394.86 | 4.52 | 3.31 |
We’ve added TTFT benchmarks to torchAO and compared different quantization + sparsity speedups for prefill / decoding. During prefill, we are compute bound and find that dynamic quantization offers greater speedups over weight-only quantization, which is faster for prefill. We’ve also added an option for int8 dynamic quantization that will selectively use prefill during LLM decoding.
The use_fp8_all_gather_only was an experimental flag, off by default, which was not marketed and not used by anyone as far as we know. We are removing it to simplify the code.
Before
config = Float8LinearConfig(
...,
# the option below is being removed
use_fp8_all_gather_only = True,
)
convert_to_float8_training(model, config=config, ...)
After
The use_fp8_all_gather_only option is no longer supported.
Full Changelog: https://github.com/pytorch/ao/compare/v0.7.0...v0.8.0-rc2
Nothing published for this version
We are excited to announce the 0.6.1 release of torchao! This release adds support for Auto-Round support, Float8 Axiswise scaled training, a BitNet t
We are excited to announce the 0.6.1 release of torchao! This release adds support for Auto-Round support, Float8 Axiswise scaled training, a BitNet training recipe, an implementation of AWQ and much more!
Auto-Round is a new weight-only quantization algorithm, it has as achieved superior accuracy compared to GPTQ, AWQ, and OmniQuant across 11 tasks, particularly excelling in low-bit quantization (e.g., 2-bits and 3-bits). Auto-Round supports quantization from 2 to 8 bits, involves low tuning costs, and imposes no additional overhead during inference. Key results are summarized below, with detailed information available in our paper, GitHub repository, and Hugging Face low-bit quantization leaderboard.
from torchao.prototype.autoround.core import prepare_model_for_applying_auto_round_
from torchao.prototype.autoround.core import apply_auto_round
prepare_model_for_applying_auto_round_(
model,
is_target_module=is_target_module,
bits=4,
group_size=128,
iters=200,
device=device,
)
input_ids_lst = []
for data in dataloader:
input_ids_lst.append(data["input_ids"].to(model_device))
multi_t_input_ids = MultiTensor(input_ids_lst)
out = model(multi_t_input_ids)
quantize_(model, apply_auto_round(), is_target_module)
We added experimental support for rowwise scaled float8 gemm to torchao.float8, with per-gemm-input configurability to enable exploration of various recipes. Here is how a user can configure all-axiswise scaling
# all-axiswise scaling
config = torchao.float8.config.recipe_name_to_linear_config(Float8LinearRecipeName.ALL_AXISWISE)
m = torchao.float8.convert_to_float8_training(config)
# or, a custom recipe by @lw where grad_weight is left in bfloat16
config = torchao.float8.config.recipe_name_to_linear_config(Float8LinearRecipeName.LW_AXISWISE_WITH_GW_HP)
m = torchao.float8.convert_to_float8_training(config)
Early performance benchmarks show all-axiswise scaling achieve a 1.13x speedup vs bf16 on torchtitan / LLaMa 3 8B / 8 H100 GPUs (compared to 1.17x from all-tensorwise scaling in the same setup), and loss curves which match to bf16 and all-tensorwise scaling. Further performance and accuracy benchmarks will follow in future releases.
Adds recipe for doing BitNet b1.58](https://arxiv.org/abs/2402.17764) ternary weights clamping.
from torchao.prototype.quantized_training import bitnet_training
from torchao import quantize_
model = ...
quantize_(model, bitnet_training())
Notably: Our implementation utilizes INT8 Tensor Cores to make up for this loss in speed. In fact, our implementation is faster than BF16 training in most cases.
Perplexity and performance measured on A100 GPU:
| Model | Quantization | Tokens/sec | Throughput (GB/sec) | Peak Mem (GB) | Model Size (GB) |
|---|---|---|---|---|---|
| Llama-2-7b-chat-hf | bfloat16 | 107.38 | 1418.93 | 13.88 | 13.21 |
| awq-hqq-int4 | 196.6 | 761.2 | 5.05 | 3.87 | |
| awq-uint4 | 43.59 | 194.93 | 7.31 | 4.47 | |
| int4wo-hqq | 209.19 | 804.32 | 4.89 | 3.84 | |
| int4wo-64 | 201.14 | 751.42 | 4.87 | 3.74 |
Usage:
from torchao.prototype.awq import insert_awq_observer_, awq_uintx, AWQObservedLinear
quant_dtype = torch.uint4
group_size = 64
calibration_limit = 10
calibration_seq_length = 1024
model=model.to(device)
insert_awq_observer_(model,calibration_limit, calibration_seq_length, quant_dtype=quant_dtype, group_size=group_size)
with torch.no_grad():
for batch in calibration_data:
model(batch.to(device))
is_observed_linear = lambda m, fqn: isinstance(m, AWQObservedLinear)
quantize_(model, awq_uintx(quant_dtype=quant_dtype, group_size = group_size), is_observed_linear)
weights_only=True for torch.load (#630)On NVIDIA GPUs, INT8 Tensor Cores is approximately 2x faster than their BF16/FP16 counterparts. In mixed-precision training, we can down-cast activations and weights dynamically to INT8 to leverage faster matmuls. However, since INT8 has very limited range [-128,127], we perform row-wise quantization, similar to how INT8 post-training quantization (PTQ) is done. Weight is still in original precision.
from torchao.prototype.quantized_training import int8_mixed_precision_training, Int8MixedPrecisionTrainingConfig
from torchao.quantization import quantize_
model = ...
# apply INT8 matmul to all 3 matmuls
quantize_(model, int8_mixed_precision_training())
# customize which matmul is left in original precision.
config = Int8MixedPrecisionTrainingConfig(
output=True,
grad_input=True,
grad_weight=False,
)
quantize_(model, int8_mixed_precision_training(config))
End2end speed benchmark using benchmarks/quantized_training/pretrain_llama2.py
| Model & GPU | bs x seq_len | Config | Tok/s | Peak mem (GB) |
|---|---|---|---|---|
| Llama2-7B, A100 | 8 x 2048 | BF16 (baseline) | ~4400 | 59.69 |
| Llama2-7B, A100 | 8 x 2048 | INT8 mixed-precision | ~6100 (+39%) | 58.28 |
| Llama2-1B, 4090 | 16 x 2048 | BF16 (baseline) | ~17,900 | 18.23 |
| Llama2-1B, 4090 | 16 x 2048 | INT8 mixed-precision | ~30,700 (+72%) | 18.34 |
fpx to floatx (#877)No significant security updates in this release.
Full Changelog: https://github.com/pytorch/ao/compare/v0.5.0...v0.6.1
We are excited to announce the 0.5 release of torchao! This release adds support for memory efficient inference, float8 training and inference, int8 q
We are excited to announce the 0.5 release of torchao! This release adds support for memory efficient inference, float8 training and inference, int8 quantized training, HQQ, automatic mixed-precision quantization through bayesian optimization, sparse marlin, and integrations with HuggingFace, SGLang, and diffusers.
We've added support for Llama 3.1 to the llama benchmarks in TorchAO and added new features and improvements as a proof of concept for memory efficient inference. These additions allow us to to do 130k context length inference with Llama 3.1-8B with only 18.91 GB memory if we combine with kv cache quantization, int4 weight only quantization and linear causal mask.
General savings depend on technique and context length as can be seen in the following graph:
torchao.float8 implements training recipes with the scaled float8 dtypes, as laid out in https://arxiv.org/abs/2209.05433.
With torch.compile on, current results show throughput speedups of up to 1.5x on 128 H100 GPU LLaMa 3 70B pretraining jobs (details)
from torchao.float8 import convert_to_float8_training
convert_to_float8_training(m, module_filter_fn=...)
And for an end-to-minimal training recipe of pretraining with float8, you can check out torchtitan.
We have introduced two new quantization APIs for Float8 inference:
Float8 Weight-Only Quantization: A new quant_api float8_weight_only() has been added to apply float8 weight-only symmetric per-channel quantization to linear layers.
Float8 Dynamic Activation and Weight Quantization: A new quant_api float8_dynamic_activation_float8_weight() has been introduced to apply float8 dynamic symmetric quantization to both activations and weights of linear layers. By default PerTensor scaling. We have also added an option to do PerRow scaling of both activations and weights. By computing scales at a finer granularity, it can potentially reduce the overall quantization error and increase performance by reducing dynamic quantization overhead.
Example usage:
import torch
from torchao.quantization import quantize_, float8_weight_only, float8_dynamic_activation_float8_weight, PerRow
# Create a model
model = YourModel()
# Apply float8 weight-only quantization
quantize_(model, float8_weight_only())
# Apply float8 dynamic activation and weight quantization
quantize_(model, float8_dynamic_activation_float8_weight())
# Apply PerRow scaling to weight and activations
quantize_(linear_module, float8_dynamic_activation_float8_weight(granularity=PerRow()))
Notes:
float8_dynamic_activation_float8_weight requires CUDA devices with compute capability 8.9 or higher for hardware acceleration.@gau-nernst introduced 2 experimental works on training using INT8.
from torchao.quantization import quantize_
from torchao.prototype.quantized_training import int8_weight_only_quantized_training, int8_mixed_precision_training
model = YourModel()
# apply INT8 quantized training
quantize_(model, int8_weight_only_quantized_training())
# apply INT8 mixed-precision training
quantize_(model, int8_mixed_precision_training())
For more information and benchmark results, see README and the respective PR (#644 and #748)
hqq is added to existing torchao APIs, it gives improvements on model accuracy and leverages the existing efficient kernels in torchao. We enabled hqq for int4_weight_only API:
quantize_(model, int4_weight_only(group_size, use_hqq=True)
We also added this to the uintx api for accuracy experiments (current uintx kernels are slow):
quantize_(model, uintx_weight_only(torch.uint2, group_size, use_hqq=True)
We provided a Bayesian Optimization (BO) tool leveraging Ax to auto search mixed-precision weight-only quantization configuration, i.e., bit width and group size of intN_weight_only(bit_width, group_size) for each layer. It also includes a sensitivity analysis tool to calculate layer-wise average Hessian trace and average fisher information matrix trace, which is an optional step to customize and improve BO search.
To optimize for model accuracy under a model size constraint (GB):
python --BO_acc_modelsize.py --checkpoint=/tmp/Meta-Llama-3-8B --model_size_constraint=6.0
To optimize for inference throughput under a model perplexity constraint:
python --BO_acc_throughput.py --checkpoint=/tmp/Meta-Llama-3-8B --ppl_constraint=7.5
For more detailed usage, please refer to this README. The mixed-precision quantization searched by this tool reduces 20.1% model size with 2.8% perplexity reduction, and improves 15.1% inference throughput with 3.2% perplexity reduction on the Llama3-8B model compared to int8 uniform quantization.
@Diogo-V added sparse-marlin, a W4AFP16 2:4 sparse kernel, support to TorchAO. On Meta LLama3, we observe a 25% tok/s increase (180 -> 226) compared to our existing int4-wo implementation.
from torchao.quantization.quant_api import quantize_, int4_weight_only
from torchao.dtypes import MarlinSparseLayoutType
quantize_(model, int4_weight_only(layout_type=MarlinSparseLayoutType()))
| Model | Technique | Tokens/Second | Memory Bandwidth (GB/s) | Peak Memory (GB) | Model Size (GB) |
|---|---|---|---|---|---|
| Llama-3-8B | Base (bfloat16) | 95.64 | 1435.54 | 16.43 | 15.01 |
| int8dq | 8.61 | 64.75 | 9.24 | 7.52 | |
| int8wo | 153.03 | 1150.80 | 10.42 | 7.52 | |
| int4wo-64 | 180.80 | 763.33 | 6.88 | 4.22 | |
| int4wo-64-sparse-marlin | 226.02 | 689.20 | 5.32 | 3.05 |
torchao is integrated into huggingface: https://huggingface.co/docs/transformers/main/en/quantization/torchao now you can use int4_weight_only, int8_weight_only and int8_dynamic_activation_int8_weight through TorchAoConfig in huggingface. Currently available in huggingface main branch only.
torchao is also integrated into sglang (https://github.com/sgl-project/sglang/pull/1341) for llama3 model, you can try out with:
python3 -m sglang.bench_latency --model meta-llama/Meta-Llama-3-8B --batch-size 1 --input 128 --output 8 --torchao-config int4wo-128
Supported configurations are ["int4wo-<group_size>", "int8wo", "int8dq", "fp8wo" (only available in torchao 0.5+)]
diffusers-torchao provides end-to-end inference and experimental training recipes to use torchao with diffusers in this repo. We demonstrate 53.88% speedup on Flux.1-Dev* and 27.33% speedup on CogVideoX-5b when comparing compiled quantized models against their standard bf16 counterparts.
# torchao 0.4.0
from torchao.quantization import quantize_, int4_weight_only
quantize_(my_model, int4_weight_only(inner_k_tiles=8))
# torchao 0.5.0
from torchao.quantization import quantize, int4_weight_only
quantize_(my_model, int4_weight_only(layout_type=TensorCoreTiledLayoutType(inner_k_tiles=8)))
We refactored QAT to use tensor subclasses instead of module swap. This works well with torchtune and FSDP2, but currently lacks support for FSDP1 and DDP. As a fallback for these distribution strategies, please continue to use the old module swap flows.
# torchao 0.4.0: This uses the module swap flow
# torch 0.5.0 + FSDP2: This uses the tensor subclass flow
from torchao.quantization.prototype.qat import Int8DynActInt4WeightQATQuantizer
quantizer = Int8DynActInt4WeightQATQuantizer()
model = quantizer.prepare(model)
train(model)
model = quantizer.convert(model)
# torchao 0.5.0 + DDP or FSDP1: This uses the module swap flow
from torchao.quantization.prototype.qat._module_swap_api import Int8DynActInt4WeightQATQuantizerModuleSwap
quantizer = Int8DynActInt4WeightQATQuantizerModuleSwap()
model = quantizer.prepare(model)
train(model)
model = quantizer.convert(model)
scripts to torchao/_models https://github.com/pytorch/ao/pull/591weights_only=True https://github.com/pytorch/ao/pull/630_quantized_linear for better extensibility https://github.com/pytorch/ao/pull/634device before quantization https://github.com/pytorch/ao/pull/699to(device=device_name) for Uintx https://github.com/pytorch/ao/pull/722torch.float32 in _foreach_norm https://github.com/pytorch/ao/pull/727torch.uint1 to torch.uint7 for Uintx tensor subclass https://github.com/pytorch/ao/pull/672CPUOffloadOptimizer default https://github.com/pytorch/ao/pull/742torch.compiler.is_compiling https://github.com/pytorch/ao/pull/739.to(device) op https://github.com/pytorch/ao/pull/595local_scale_tensor to fp32 for precompute of float8 dynamic scaling https://github.com/pytorch/ao/pull/713sparsify_ https://github.com/pytorch/ao/pull/741We were able to close about 70% of tasks for 0.5.0, which will now spill over into upcoming releases. We will post a list for 0.6.0 next, which we aim to release at the end of September 2024. We want to follow a monthly release cadence until further notice.
Full Changelog: https://github.com/pytorch/ao/compare/v0.4.0...v0.5.0-rc1
apply_sparse_semi_structured has been deprecated in favor of sparsify_ which matches the quantize_ API https://github.com/pytorch/ao/pull/473 ``` pyth…
We are excited to announce the 0.4 release of torchao! This release adds support for KV cache quantization, quantization aware training (QAT), low bit optimizer support, composing quantization and sparsity, and more!
We've added support for KV cache quantization, showing a peak memory reduction from 19.7 -> 19.2 GB on Llama3-8B at an 8192 context length. We plan to investigate Llama3.1 next.
<img src="https://github.com/user-attachments/assets/31946f46-e8eb-45c2-ac1c-3a7d981c58a2" width="300" height="auto">
We now support two QAT schemes for linear layers: Int8 per token dynamic activations + int4 per group weights, and int4 per group weights (using the efficient tinygemm int4 kernel after training). Users can access this feature by transforming their models before and after training using the appropriate quantizer, for example:
from torchao.quantization.prototype.qat import Int8DynActInt4WeightQATQuantizer
# Quantizer for int8 dynamic per token activations +
# int4 grouped per channel weights, only for linear layers
qat_quantizer = Int8DynActInt4WeightQATQuantizer()
# Insert "fake quantize" operations into linear layers.
# These operations simulate quantization numerics during
# training without performing any dtype casting
model = qat_quantizer.prepare(model)
# Convert fake quantize to actual quantize operations
model = qat_quantizer.convert(model)
Initial evaluation results indicate that QAT in torchao can recover up to 96% of quantized accuracy degradation on hellaswag and up to 68% of quantized perplexity degradation on wikitext for Llama3 compared to post-training quantization (PTQ). For more details, please refer to the README and this blog post.
We've added support for composing int8 dynamic quantization with 2:4 sparsity, using the quantize_ API. We also added SAM benchmarks that show a 7% speedup over standalone sparsity / int8 dynamic quantization here.
from torchao.quantization import quantize_, int8_dynamic_activation_int8_semi_sparse_weight
quantize_(model, int8_dynamic_activation_int8_semi_sparse_weight())
@gau-nernst added implementations for 4-bit, 8-bit, and FP8 Adam with FSDP2/FSDP support. Our API is a drop-in replacement for torch.optim.Adam and can be used as follows:
from torchao.prototype.low_bit_optim import Adam8bit, Adam4bit, AdamFp8
from torchao.prototype.low_bit_optim import AdamW8bit, AdamW4bit, AdamWFp8
model = ...
optim = Adam8bit(model.parameters()) # replace with Adam4bit and AdamFp8 for the 4 / fp8 versions
For more information about low bit optimizer support please refer to our README.
@bdhirsh @jeromeku @yanbing-j @manuelcandales @larryliu0820 added torch.compile support for NF4 Tensor, custom CUDA int4 tinygemm unpacking ops, and several bugfixes to torchao
quantize has been renamed to quantize_ https://github.com/pytorch/ao/pull/467# for torchao 0.4
from torchao.quantization import quantize_, int8_weight_only
quantize_(model, int8_weight_only())
# for torchao 0.3
from torchao.quantization import quantize, int8_weight_only
quantize(model, int8_weight_only())
apply_sparse_semi_structured has been deprecated in favor of sparsify_ which matches the quantize_ API https://github.com/pytorch/ao/pull/473# for torchao 0.4
from torchao.sparsity import _sparsify, semi_sparse_weight
sparsify_(model, semi_sparse_weight())
# for torchao 0.3
from torchao.sparsity import apply_sparse_semi_structured
apply_sparse_semi_structured(model)
torchao.float8, enabling float8 training support https://github.com/pytorch/ao/pull/551 https://github.com/pytorch/ao/pull/529tinygemm unpacking ops https://github.com/pytorch/ao/pull/415 Int8DynActInt4WeightQuantizer and Int4WeightOnlyQuantizer https://github.com/pytorch/ao/pull/475 https://github.com/pytorch/ao/pull/479model.to for int4/int8 weight only quantized models https://github.com/pytorch/ao/pull/486 https://github.com/pytorch/ao/pull/522TensorCoreTiledAQTLayout https://github.com/pytorch/ao/pull/520fake_quantize_affine op with mask support https://github.com/pytorch/ao/pull/492 https://github.com/pytorch/ao/pull/500fake_quantize_affine primitive https://github.com/pytorch/ao/pull/527unwrap_tensor_subclass https://github.com/pytorch/ao/pull/595TORCH_VERSION_AFTER_* https://github.com/pytorch/ao/pull/433torch.compile support for NF4Tensor https://github.com/pytorch/ao/pull/544int4pack_mm error https://github.com/pytorch/ao/pull/517segment-anything-fast benchmarks for composed quantization + sparsity https://github.com/pytorch/ao/pull/457torchao no long depends on torch https://github.com/pytorch/ao/pull/449benchmark_model now accepts args and kwargs and supports cpu and mps backends https://github.com/pytorch/ao/pull/586 https://github.com/pytorch/ao/pull/406Quantizer now uses logging instead of print https://github.com/pytorch/ao/pull/472_replace_linear_8da4w https://github.com/pytorch/ao/pull/451LinearActQuantizedTensor https://github.com/pytorch/ao/pull/542Full Changelog: https://github.com/pytorch/ao/compare/v0.3.1-rc1...v0.4.0-rc1
We were able to close about 60% of tasks for 0.4.0, which will now spill over into upcoming releases. We will post a list for 0.5.0 next, which we aim to release at the end of August 2024. We want to follow a monthly release cadence until further notice.
Nothing published for this version
Deprecate top level quantization APIs https://github.com/pytorch/ao/pull/344
We are excited to announce the 0.3 release of torchao! This release adds support for a new quantize API, MX format, FP6 dtype and bitpacking, 2:4 sparse accelerated training and benchmarking infra for llama2/llama3 models.
quantize API (https://github.com/pytorch/ao/pull/256)We added a tensor subclass based quantization API, see docs and README for details on usage, this is planned to replace all existing quantization APIs in torchao for torch 2.4 and later.
You can now accelerate training with 2:4 sparsity, using the runtime pruning + compression kernels written by xFormers. These kernels process a 4x4 sub-tile to be 2:4 sparse in both directions, to handle both the forward and backward pass when training. We see a 1.3x speedup for the MLP layers of ViT-L across a forward and backwards pass.
We added prototype support for MX format for training and inference with a reference native PyTorch implementation of training and inference primitives for using MX accelerated matrix multiplications. The MX numerical formats are new low precision formats with recent acceptance into the OCP spec: https://www.opencompute.org/documents/ocp-microscaling-formats-mx-v1-0-spec-final-pdf
We added a stable way to benchmark llama2 and llama3 models that includes perf/accuracy comparisons. See torchao/_models/llama/benchmarks.sh for more details.
@gau-nernst Added support for FP6 dtype and mixed matmul FP16 x FP6 kernel with support for torch.compile. Benchmark results show a 2.3x speedup over BF16 baseline for meta-llama/Llama-2-7b-chat-hf
@vayuda, @melvinebenezer @CoffeeVampir3 @andreaskoepf Added support for packing/unpacking lower bit dtypes leveraging torch.compile to generate the kernels for this and added UInt2 and Bitnet tensor based on this approach.
Added the kernel written by @AdnanHoque to torchao with speedups compared to the cuBLAS kernel for batch size <=16
apply_weight_only_int8_quant(model) or change_linear_weights_to_int8_woqtensors(model)
-->
# for torch 2.4+
from torchao.quantization import quantize, int8_weight_only
quantize(model, int8_weight_only())
# for torch 2.2.2 and 2.3
from torchao.quantization.quant_api import change_linear_weights_to_int8_woqtensors
change_linear_weights_to_int8_woqtensors(model)
apply_dynamic_quant(model) or change_linear_weights_to_int8_dqtensors(model)
-->
# Fuse the int8*int8 -> int32 matmul and subsequent mul op avoiding materialization of the int32 intermediary tensor
torch._inductor.config.force_fuse_int_mm_with_mul = True
# for torch 2.4+
from torchao.quantization import quantize, int8_dynamic_activation_int8_weight
quantize(model, int8_dynamic_activation_int8_weight())
# for torch 2.2.2 and 2.3
from torchao.quantization.quant_api import change_linear_weights_to_int8_dqtensors
change_linear_weights_to_int8_dqtensors(model)
change_linear_weights_to_int4_wotensors(model)
-->
# for torch 2.4+
from torchao.quantization import quantize, int4_weight_only
quantize(model, int4_weight_only())
# for torch 2.2.2 and 2.3
from torchao.quantization.quant_api import change_linear_weights_to_int4_woqtensors
change_linear_weights_to_int4_woqtensors(model)
quantize https://github.com/pytorch/ao/pull/256AQTLayout, PlainAQTLayout and TensorCoreTiledAQTLayout https://github.com/pytorch/ao/pull/278quantize https://github.com/pytorch/ao/pull/294quantization_factor https://github.com/pytorch/ao/pull/298quantize https://github.com/pytorch/ao/pull/301register_apply_tensor_subclass https://github.com/pytorch/ao/pull/366NF4Tensors https://github.com/pytorch/ao/pull/297hf_eval.py https://github.com/pytorch/ao/pull/341AffineQuantizedTensor based workflow doc and examples https://github.com/pytorch/ao/pull/277AUTOQUANT_CACHE docs for reusing the same quantization plan https://github.com/pytorch/ao/pull/329quantize to doc page https://github.com/pytorch/ao/pull/367AffineQuantizedTensor to torchao/dtypes https://github.com/pytorch/ao/pull/272HQQ triton_mm test https://github.com/pytorch/ao/pull/306Full Changelog: https://github.com/pytorch/ao/compare/v0.2.0...v0.3.0-rc1
We were able to close about 60% of tasks for 0.3.0, which will now spill over into upcoming releases. We will post a list for 0.4.0 next, which we aim to release at the end of July 2024. We want to follow a monthly release cadence until further notice.
EDIT: We made a patch release for 0.3.1 to include 2 more PRs so now ao has no runtime dependencies https://github.com/pytorch/ao/pull/449 and https://github.com/pytorch/ao/pull/455
PyTorch core has recently shipped a new custom op registration mechanism with torch.library with the benefit being that custom ops will compose with a
PyTorch core has recently shipped a new custom op registration mechanism with torch.library with the benefit being that custom ops will compose with as many PyTorch subsystems as possible most notably NOT graph breaking with torch.compile()
We'd added some documentation for how you could register your own custom ops https://github.com/pytorch/ao/tree/main/torchao/csrc and if you learn better via example you can follow this PR https://github.com/pytorch/ao/pull/135 to add your own custom ops to torchao.
Most notably these instructions were leveraged by @gau-nernst to integrate some new custom ops for fp6 support https://github.com/pytorch/ao/pull/223
One key benefit of integrating your kernels in torchao directly is we thanks to our manylinux GPU support can ensure that CPU/CUDA kernels that you've added will work on as many devices and cuda versions as possible https://github.com/pytorch/ao/pull/176
@jeromeku was our community champion merging support for
@gau-nernst merged fp6 support showing up to 8x speedups on an fp16 baseline for small batch size inference https://github.com/pytorch/ao/pull/223
@weifengpy merged support for composing FSDP2 with NF4 which makes it easy to implement algorithms like QLoRA + FSDP without writing any CUDA or C++ code. This work also provides a blueprint for how to compose smaller dtypes with FSDP https://github.com/pytorch/ao/pull/150 most notably by implementing torch.chunk(). We hope the broader community uses this work to experiment more heavily at the intersection of distributed and quantization research and inspires many more studies such as the ones done by Answer.ai https://www.answer.ai/posts/2024-03-06-fsdp-qlora.html
Int4WeightOnlyQuantizer (https://github.com/pytorch/ao/pull/119)is_gpt_fast specialization from GTPQ (https://github.com/pytorch/ao/pull/172)Int8DynActInt4WeightLinear module swap (https://github.com/pytorch/ao/pull/151)NF4Tensor.to to use device kwarg (https://github.com/pytorch/ao/pull/158)quantize_activation_per_token_absmax perf regression (https://github.com/pytorch/ao/pull/253)Full Changelog: https://github.com/pytorch/ao/compare/v0.2.0...v0.2.1
We were able to close about half of tasks for 0.2.0, which will now spill over into upcoming releases. We will post a list for 0.3.0 next, which we aim to release at the end of May 2024. We want to follow a monthly release cadence until further notice.
We’re excited to announce the release of TorchAO v0.1.0! TorchAO is a repository to host architecture optimization techniques such as quantization and
We’re excited to announce the release of TorchAO v0.1.0! TorchAO is a repository to host architecture optimization techniques such as quantization and sparsity and performance kernels on different backends such as CUDA and CPU. In this release, we added support for a few quantization techniques like int4 weight only GPTQ quantization, added nf4 dtype support for QLoRA and sparsity features like WandaSparsifier, we also added autotuner that can tune triton integer matrix multiplication kernels on cuda.
Note: TorchAO is currently in a pre-release state and under extensive development. The public APIs should not be considered stable. But we welcome you to try out our APIs and offerings and provide any feedback on your experience.
torchao 0.1.0 will be compatible with PyTorch 2.2.2 and 2.3.0, ExecuTorch 0.2.0 and TorchTune 0.1.0.
change_linear_weights_to_int8_dqtensors, change_linear_weights_to_int8_woqtensors and change_linear_weights_to_int4_woqtensors (#1)apply_weight_only_int8_quant and apply_dynamic_quant (#1)Int4WeightOnlyQuantizer and Int4WeightOnlyGPTQQuantizer used in TorchTune (#119, #116)Int8DynActInt4WeightQuantizer and Int8DynActInt4WeightGPTQQuantizer, used in ExecuTorch (#74) (available after torch 2.3.0 and later)WandaSparsifier that prunes both weights and activations (#22)autotuner for int mm Triton kernels (#41)nf4 tensor subclass and nf4 linear (#37, #40, #62)uint4 dtype tensor subclass (#13)torchao-nightly release (#54)nf4 tensor (#54)Int8DynActInt4WeightGPTQQuantizeruint4 tensor subclass is going to be merged into pytorch core in the futurequant_primitives.py will be deduplicated with similar quantize/dequantize ops in PyTorch laterNothing published for this version
Nothing published for this version
Your coding agent can read these notes before it upgrades. Set up the MCP server →