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拡張がプロジェクトに含まれています。
対象ユーザー
低精度量子化と数値精度の探索に注力するデータサイエンティストおよびAI研究者向けです。
主な特徴
- 広範なフォーマット対応:FP8(e4m3、e5m2)、FP4(e2m1)、INT8など、さまざまなMX互換フォーマットをサポート。
- 柔軟な設定:
mx_specs辞書を使用して、スケールビット、重みと活性化の要素フォーマット、ブロックサイズを構成可能。 - シームレスな統合:手動でのPyTorchモジュールの置き換え、または
mx_mapping.inject_pyt_opsを通じた操作の自動挿入の2つの統合パスを提供。 - 高いパフォーマンス:PyTorch GPUの既知の数値不正確さを回避し、シミュレーション速度を向上させるカスタムCUDAカーネルを含む。
- 包括的なカバー範囲:前方伝搬と逆伝搬の量子化、GELU、Softmax、LayerNormなどの要素演算をカバー。
関連
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト