pytorch/ao
PyTorch native quantization and sparsity for training and inference
What it solves
TorchAO 提供一個原生 PyTorch 函式庫,用於優化 AI 模型,使其更快且更省記憶體。它解決了模型大小與準確度之間的常見權衡,讓使用者能在不顯著降低品質的前提下,減少大型模型(如 LLM 與 diffusion 模型)的記憶體佔用,並加速訓練與推論。
How it works
TorchAO 實作了多種架構優化技術:
- Quantization:將模型權重與激活轉換為較低精度的格式(如 int4、int8、float8),降低記憶體使用並提升吞吐量。
- Quantization-Aware Training (QAT):為避免量化時的精度損失,允許模型在訓練過程中適應低精度。
- Sparsity:使用 2:4 半結構化稀疏化移除冗餘權重,進一步提升速度。
- uma-native integration:與
torch.compile()與FSDP2無縫結合,於 CUDA、XPU、CPU、ARM 等多種硬體上提供高效執行。
Who it’s for
此函式庫設計給需要在受限硬體上部署大規模模型、加速大模型預訓練,或透過 ExecuTorch 為邊緣裝置優化模型的 AI 研究員與工程師。
Highlights
- Training Speedups:使用 float8 訓練可讓 Llama-3.1-70B 的預訓練速度提升至 1.5 倍。
- Inference Gains:將 Llama-3-8B 量化為 int4 可達到 1.89 倍更快的推論速度,且記憶體使用減少 58%。
- Broad Integration:內建支援 Hugging Face Transformers、Diffusers、vLLM 與 SGLang。
- Memory Efficiency:提供量化優化器(AdamW 4/8-bit)與 CPU offloading,將 VRAM 需求降低最高可達 60%。