pytorch/ao · 上手攻略
- 仓库:pytorch/ao
- 链接:https://github.com/pytorch/ao
- 分类:llm-infra / 模型量化与稀疏化
- 作者:spark
- 更新:2026-07-27
是什么
torchao(即 pytorch/ao 仓库发布的 PyPI 包名)是 PyTorch 官方主推的 原生量化与稀疏化库,主打「训练-推理一致」「一行代码改动即可接入」,覆盖从权重量化(int4/int8/fp8/NF4)、激活量化、2:4 稀疏、低比特 optimizer(4bit/8bit/fp8 AdamW)、量化感知训练(QAT),到 float8 预训练、CPU offload 的全链路。仓库 2024 年 9 月由 Meta/PyTorch 团队正式对外发布,arXiv 2507.16099(2025 TorchAO 论文)已经在 CodeML @ ICML 2025 接收。
它不是一个独立训练框架,而是 PyTorch 模型对象上的 transform:通过 quantize_(model, config) 或在 Hugging Face from_pretrained 阶段挂上 TorchAoConfig 来完成。同一个量化模型可以无缝跑在 CUDA、XPU、CPU 上,并且被 Unsloth、Hugging Face Transformers / Diffusers / PEFT、vLLM、SGLang、Axolotl、TorchTune、TorchTitan、ExecuTorch 等主流生态直接集成。
解决什么问题
- LLM/扩散模型推理显存爆炸:把 Llama-3-8B 压成 int4 可以少用 58% 显存、推理 1.89× 加速;Qwen3-4B 在 iPhone 15 Pro 上跑出 14.8 tok/s 只占 3.4GB。
- 大模型预训练太慢:float8 rowwise 训练在 Crusoe 2K H200 上拿到 1.34–1.43× 加速;MXFP8 训练在 B200 上对 Llama-4 Scout 的 MoE 层拿到约 1.45× 加速,数值表现接近 bfloat16。
- PTQ 精度掉太多:QAT 在 Llama-3-8B 上能恢复 hellaswag 96% 的精度损失、wikitext 上 68% 的 PPL 损失,可与 LoRA 结合再提速 1.89×。
- AdamW 优化器吃显存:用
AdamW8bit/AdamW4bit/AdamWFp8把优化器状态内存直接砍掉 2–4×。
快速安装
最新稳定版(推荐):
pip install torchao
不同 CUDA 版本对应的 wheel:
# CUDA 12.6 / 12.9 / 仅 CPU / Intel XPU / nightly
pip install torchao --index-url https://download.pytorch.org/whl/cu126
pip install torchao --index-url https://download.pytorch.org/whl/cu129
pip install torchao --index-url https://download.pytorch.org/whl/cpu
pip install torchao --index-url https://download.pytorch.org/whl/xpu
pip install --pre torchao --index-url https://download.pytorch.org/whl/nightly/cu128
可选的 SOTA 内核依赖 MSLK(稳定 1.0.0 / nightly 与 torchao 同步):
pip install mslk-cuda==1.0.0
pip install --pre mslk --index-url https://download.pytorch.org/whl/nightly/cu128
源码开发模式(必须 --no-build-isolation):
USE_CUDA=1 pip install -e . --no-build-isolation
注意:具体 CUDA / cuDNN / PyTorch 兼容性请查仓库 Issue #2919「torchao compatibility table」对照矩阵;CUDA 13 + 最新 nightly 的组合每月都在变。
核心用法
1. 一行量化(quantize_ API,适合任意 nn.Module)
import torch
from torchao.quantization import Int4WeightOnlyConfig, quantize_
model = ... # 你的 nn.Module
if torch.cuda.is_available():
quantize_(model, Int4WeightOnlyConfig(
group_size=32,
int4_packing_format="tile_packed_to_4d",
int4_choose_qparams_algorithm="hqq",
))
elif torch.xpu.is_available():
quantize_(model, Int4WeightOnlyConfig(
group_size=32,
int4_packing_format="plain_int32",
))
2. Hugging Face Transformers 集成(推荐给 LLM)
from transformers import TorchAoConfig, AutoModelForCausalLM
from torchao.quantization import Float8DynamicActivationFloat8WeightConfig
from torchao.quantization.granularity import PerRow
quantization_config = TorchAoConfig(
quant_type=Float8DynamicActivationFloat8WeightConfig(granularity=PerRow())
)
quantized_model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen3-32B",
dtype="auto",
device_map="auto",
quantization_config=quantization_config,
)
加载完即可用 vLLM 起服务:
VLLM_DISABLE_COMPILE_CACHE=1 vllm serve pytorch/Qwen3-32B-FP8 \
--tokenizer Qwen/Qwen3-32B -O3
3. Hugging Face Diffusers 集成(扩散模型)
import torch
from diffusers import DiffusionPipeline, PipelineQuantizationConfig, TorchAoConfig
from torchao.quantization import Int8WeightOnlyConfig
from torchao.quantization.granularity import PerGroup
pipeline = DiffusionPipeline.from_pretrained(
"black-forest-labs/FLUX.1-dev",
quantization_config=PipelineQuantizationConfig(
quant_mapping={"transformer": TorchAoConfig(
Int8WeightOnlyConfig(granularity=PerGroup(128))
)}
),
torch_dtype=torch.bfloat16,
device_map="cuda",
)
Flux.1-Dev + CogVideoX-5b 在 H100 上常见 1.54× / 1.27× 加速。
4. 量化感知训练 QAT(精度掉太多时)
from torchao.quantization import quantize_, Int8DynamicActivationIntxWeightConfig, PerGroup
from torchao.quantization.qat import QATConfig
base_config = Int8DynamicActivationIntxWeightConfig(
weight_dtype=torch.int4,
weight_granularity=PerGroup(32),
)
quantize_(my_model, QATConfig(base_config, step="prepare"))
# ... 正常训练若干 step ...
quantize_(my_model, QATConfig(base_config, step="convert"))
Unsloth、Axolotl、TorchTune 都已经把 QAT 写成 recipe,开箱即用;QAT + LoRA 比 vanilla QAT 快 1.89×。
5. Float8 预训练(需要 ≥2K GPU 规模才划算)
from torchao.float8 import convert_to_float8_training
convert_to_float8_training(m)
通常与 TorchTitan + FSDP2 配合使用,2025 年在 Llama-3.1-70B/405B 上拿到 1.43–1.51× 预训练加速。
6. 低比特优化器(节省 AdamW 显存)
from torchao.optim import AdamW8bit, AdamW4bit, AdamWFp8, CPUOffloadOptimizer
optim = AdamW8bit(model.parameters()) # 或 AdamW4bit / AdamWFp8
# 单 GPU 想再省点,可以把梯度+状态 offload 到 CPU:
optim = CPUOffloadOptimizer(model.parameters(), torch.optim.AdamW, fused=True)
optim.load_state_dict(ckpt["optim"])
7. 2:4 半结构化稀疏(推理 + 训练都能加)
from torchao.sparsity.training import SemiSparseLinear, swap_linear_with_semi_sparse_linear
swap_linear_with_semi_sparse_linear(model, {"seq.0": SemiSparseLinear})
Llama-3-8B int4 + 2:4 稀疏能拿到 2.37× 吞吐、67.7% 内存下降。
典型适用场景
- 单卡/多卡部署量化 LLM:要 int4/FP8 推理、想直接用 vLLM / SGLang 服务的团队。仓库已经预量化了一批「pytorch/」前缀的 HF Hub 模型(Qwen3、Gemma-3、Llama-3 等),可以直接拉。
- 大模型预训练加速:TorchTitan + float8 rowwise / MXFP8,在 256+ H100/B200 集群上部署。
- Diffusion 模型推理加速:Flux、CogVideoX、SD 系列做 int8/fp8 量化。
- 端侧 / 移动端 LLM:通过 ExecuTorch 把 int4/int8 模型部署到 iPhone、Android、ARM CPU(仓库有 1-8 bit ARM 内核)。
- 训练时省显存:低比特 optimizer / CPU offload,单卡 405B 也敢玩。
- QAT 微调:Unsloth / Axolotl / TorchTune 都接好了,需要在低比特下保住精度的 LoRA / 全参微调场景。
坑与注意
- API 仍在快速演进:仓库月更频繁,
Int4WeightOnlyConfig、Float8DynamicActivationFloat8WeightConfig等名字、参数(int4_packing_format、granularity=PerRow())半年内会换;写脚本前一定查docs.pytorch.org/ao/main当前版本,而不是依赖本攻略里的示例长期不动。 - CUDA / PyTorch 兼容性:不要混 stable 与 nightly,stable ↔ stable、nightly ↔ nightly,否则会撞到 ABI 错。MSLK 同理。
Int8DynamicActivationIntxWeightConfig等低比特配置:PTQ 直接用会掉精度,先用 QAT 流程 prepare → 训练 → convert,再上 serving。- 量化粒度选择:weight-only 默认
PerRow()在大模型上更稳;激活量化建议用PerGroup(32/128),全PerTensor()对 LLM 精度损失较大。 quantize_在torch.compile后:torch.compile(model)之后再quantize_会触发重编译;先量化、再 compile,或干脆在quantize_(...)内传入inplace=True。- 显存评估:FP8 量化不一定省 VRAM,省的是算力(吞吐)和权重显存;激活 / KV cache 不变。要量激活得用
Float8DynamicActivationFloat8WeightConfig或Int8DynamicActivationIntxWeightConfig。 - 多卡分发:权重 dtype 改变会让
FSDP2的 bucket 切分逻辑变化,升级torchao后一定在最小集群上跑一次完整 forward+backward + checkpoint load/save。 - MXFP8 仍在 prototype:
torchao/prototype/moe_training路径下的 MXFP8 MoE 训练是实验性,prod 慎用。 - 不替代专用推理引擎:移动端走 ExecuTorch、服务端走 vLLM/SGLang,torchao 自己不带 engine。
与同类对比
| 库 | 定位 | 与 torchao 的差异 |
|---|---|---|
| bitsandbytes | int8/int4 LLM 量化 | 社区先发、生态成熟,但功能更窄(无 float8、无 QAT、无 diffusers 集成),最新版开始向 torchao 看齐。 |
| transformers optimum | HF 一站式量化 / 导出 | optimum 本身就是 torchao / ONNX / GPTQ 等多种 backend 的调度层,背后量化算子很多直接来自 torchao。 |
| AutoGPTQ / AutoAWQ | GPTQ / AWQ 权重量化 | 算法路线不同(需要校准数据),torchao 走 PTQ/QAT 不依赖校准数据集。 |
| vLLM / SGLang | LLM serving 引擎 | 自带量化 backend(vLLM 0.6+ 已把 torchao 列为官方 backend);量化 + 部署是两件事,配合使用最香。 |
| MSLK / Meta-pytorch kernels | 算子库 | torchao 的可选加速后端,提供 SOTA 内核但本身不带量化配方。 |
| Nvidia Transformer Engine | GPU 大厂量化库 | Hopper/Blackwell 上 TE 是默认选择;torchao 的优势在 PyTorch 原生 + 多硬件 + 不绑 NVIDIA。 |
一句话推荐
如果你已经在 PyTorch 生态里(HF Transformers / Diffusers / TorchTitan / ExecuTorch),先
pip install torchao,再决定要不要换库——绝大多数量化/稀疏化需求它都能少改代码搞定。
参考来源:仓库 README(github.com/pytorch/ao,2026-07-27 抓取)、TorchAO 论文 arXiv:2507.16099、docs.pytorch.org/ao/main、HF Transformers 文档 huggingface.co/docs/transformers/main/quantization/torchao。版本号 / 安装参数以官方 docs 为准——本攻略涉及的功能在 0.10+ 稳定版中均可获得,但具体 config 字段名随 minor 版本迭代。