microsoft/microxcaling
PyTorch emulation library for Microscaling (MX)-compatible data formats
解決的問題
此套件允許研究人員與資料科學家在 PyTorch 中模擬 MX 相容的資料格式與 bfloat 量化。它讓使用者能在不需依賴原生支援這些格式的專用硬體的情況下,探索不同低精度數值格式(如 FP8、FP4 和 INT8)對深度神經網路(DNN)效能與準確度的影響。
工作原理
該套件透過在較高精度(float32、bfloat16 或 fp16)中執行運算,同時將數值限制在目標 MX 或 bfloat 格式的範圍與精度內,來模擬低精度格式。它提供標準 PyTorch 模組與函數(如 torch.matmul、torch.linear 和 torch.nn.LayerNorm)的即插即用替代品。
為提升模擬速度與數值準確度,超越原生 PyTorch GPU 操作,專案包含用於量化的自訂 CUDA 延伸。
適用對象
專為專注於 DNN 中量化與數值精度探索的資料科學家與 AI 研究人員設計。
主要特色
- 廣泛格式支援:支援多種 MX 相容格式,包括 FP8(e4m3、e5m2)、FP4(e2m1)與 INT8。
- 彈性設定:使用
mx_specs字典配置比例位數、權重與激活的元素格式以及區塊大小。 - 無縫整合:提供兩種整合路徑:手動取代 PyTorch 模組,或透過
mx_mapping.inject_pyt_ops自動注入運算。 - 高效率:包含自訂 CUDA 核心,以避免已知的 PyTorch GPU 數值不準確問題並提升模擬速度。
- 全面覆蓋:涵蓋前向與反向傳播量化,以及 GELU、Softmax 與 LayerNorm 等元素運算。
相關
- 專案
- 專案
- 專案
- 專案
- 專案