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.matmultorch.lineartorch.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などの要素演算をカバー。

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト