日本語版
最新ニュース
科学&テクノロジー

deepseek-ai/deepgemm:deepgemm:きれいで効率的なfp8ジェムカーネル

DeepGEMMは、で提案されているように、きめ細かいスケーリングを備えたクリーンで効率的なFP8一般マトリックス乗算(GEMMS)のために設計されたライブラリです。 deepseek-v3。通常とミックスの両方の専門家(MOE)グループ化されたGEMMをサポートします。 Cudaで書かれたライブラリは、軽量のジャストインタイム(JIT)モジュールを使用して、実行時にすべてのカーネルをコンパイルすることにより、インストール中にコンピレーションを必要としません。 現在、DeepGEMMはNvidia Hopperテンソルコアのみをサポートしています。不正確なFP8テンソルコアの蓄積に対処するために、CUDAコアの2レベルの蓄積(昇進)を採用しています。それはからいくつかの概念を活用していますが カトラス そして かわいい、それは彼らのテンプレートや代数への大きな依存を避けます。代わりに、ライブラリはシンプルに設計されており、周囲で構成されるコアカーネル関数は1つだけです 〜300行のコード。これにより、Hopper FP8マトリックスの乗算と最適化技術を学習するためのクリーンでアクセス可能なリソースになります。 軽量のデザインにもかかわらず、DeepGemmのパフォーマンスは、さまざまなマトリックス形状にわたってエキスパートチューニングされたライブラリに一致するか、それを超えています。 NVCC 12.8のH800で、DeepSeek-V3/R1の推論(踏み出しとデコードの両方を含むが、テンソル並列症の両方を含む)で使用される可能性のあるすべての形状をテストします。すべてのスピードアップメトリックは、Cutlass 3.6に基づいて、内部および慎重に最適化された実装と比較して計算されます。 DeepGEMMはいくつかの形状ではあまりうまく動作しません。興味があれば最適化PRは歓迎されます。 密なモデルの通常のgemms m n k 計算 メモリ帯域幅 スピードアップ 64 2112 7168 206 TFLOPS 1688 GB/s。 2.7x 64 24576 1536 289 TFLOPS 2455 GB/s 1.7x…

deepseek-ai/deepgemm:deepgemm:きれいで効率的なfp8ジェムカーネル

1740539107
2025-02-26 01:02:00

DeepGEMMは、で提案されているように、きめ細かいスケーリングを備えたクリーンで効率的なFP8一般マトリックス乗算(GEMMS)のために設計されたライブラリです。 deepseek-v3。通常とミックスの両方の専門家(MOE)グループ化されたGEMMをサポートします。 Cudaで書かれたライブラリは、軽量のジャストインタイム(JIT)モジュールを使用して、実行時にすべてのカーネルをコンパイルすることにより、インストール中にコンピレーションを必要としません。

現在、DeepGEMMはNvidia Hopperテンソルコアのみをサポートしています。不正確なFP8テンソルコアの蓄積に対処するために、CUDAコアの2レベルの蓄積(昇進)を採用しています。それはからいくつかの概念を活用していますが カトラス そして かわいい、それは彼らのテンプレートや代数への大きな依存を避けます。代わりに、ライブラリはシンプルに設計されており、周囲で構成されるコアカーネル関数は1つだけです 〜300行のコード。これにより、Hopper FP8マトリックスの乗算と最適化技術を学習するためのクリーンでアクセス可能なリソースになります。

軽量のデザインにもかかわらず、DeepGemmのパフォーマンスは、さまざまなマトリックス形状にわたってエキスパートチューニングされたライブラリに一致するか、それを超えています。

NVCC 12.8のH800で、DeepSeek-V3/R1の推論(踏み出しとデコードの両方を含むが、テンソル並列症の両方を含む)で使用される可能性のあるすべての形状をテストします。すべてのスピードアップメトリックは、Cutlass 3.6に基づいて、内部および慎重に最適化された実装と比較して計算されます。

DeepGEMMはいくつかの形状ではあまりうまく動作しません。興味があれば最適化PRは歓迎されます。

密なモデルの通常のgemms

m n k 計算 メモリ帯域幅 スピードアップ
64 2112 7168 206 TFLOPS 1688 GB/s。 2.7x
64 24576 1536 289 TFLOPS 2455 GB/s 1.7x
64 32768 512 219 TFLOPS 2143 GB/s 1.8x
64 7168 16384 336 TFLOPS 2668 GB/s。 1.4x
64 4096 7168 287 TFLOPS 2320 gb/s 1.4x
64 7168 2048 295 TFLOPS 2470 GB/s 1.7x
128 2112 7168 352 TFLOPS 1509 GB/s。 2.4x
128 24576 1536 535 TFLOPS 2448 GB/s。 1.6x
128 32768 512 358 TFLOPS 2103 GB/s 1.5x
128 7168 16384 645 TFLOPS 2604 GB/s。 1.4x
128 4096 7168 533 TFLOPS 2221 GB/s 2.0x
128 7168 2048 510 TFLOPS 2277 GB/s。 1.7x
4096 2112 7168 1058 TFLOPS 527 GB/s。 1.1x
4096 24576 1536 990 TFLOPS 786 GB/s。 1.0x
4096 32768 512 590 TFLOPS 1232 GB/s。 1.0x
4096 7168 16384 1358 TFLOPS 343 GB/s 1.2x
4096 4096 7168 1304 TFLOPS 500 GB/s 1.1x
4096 7168 2048 1025 TFLOPS 697 GB/s。 1.1x

MOEモデル用のグループ化されたGEMMS(連続レイアウト)

#Groups グループあたりm n k 計算 メモリ帯域幅 スピードアップ
4 8192 4096 7168 1297 TFLOPS 418 gb/s。 1.2x
4 8192 7168 2048 1099 TFLOPS 681 GB/s 1.2x
8 4096 4096 7168 1288 TFLOPS 494 GB/s。 1.2x
8 4096 7168 2048 1093 TFLOPS 743 GB/s。 1.1x

MOEモデル用のグループ化されたGEMMS(マスクレイアウト)

#Groups グループあたりm n k 計算 メモリ帯域幅 スピードアップ
1 1024 4096 7168 1233 TFLOPS 924 GB/s。 1.2x
1 1024 7168 2048 925 TFLOPS 968 GB/s 1.2x
2 512 4096 7168 1040 TFLOPS 1288 GB/s。 1.2x
2 512 7168 2048 916 TFLOPS 1405 GB/s。 1.2x
4 256 4096 7168 932 TFLOPS 2064 GB/s。 1.1x
4 256 7168 2048 815 TFLOPS 2047 GB/s。 1.2x

  • ホッパーアーキテクチャGPU、 sm_90a サポートする必要があります
  • Python 3.8以上
  • CUDA 12.3以上
    • しかし、最高のパフォーマンスには12.8以上を強くお勧めします
  • Pytorch 2.1以上
  • Cutlass 3.6以上(Gitサブモジュールでクローン化できます)

# Submodule must be cloned
git clone --recursive [email protected]:deepseek-ai/DeepGEMM.git

# Make symbolic links for third-party (CUTLASS and CuTe) include directories
python setup.py develop

# Test JIT compilation
python tests/test_jit.py

# Test all GEMM implements (normal, contiguous-grouped and masked-grouped)
python tests/test_core.py

次に、インポートします deep_gemm あなたのPythonプロジェクトで、そして楽しんでください!

このライブラリには、GEMMカーネルのみが含まれています。 LHSスケーリング係数をTMAに合わせて転置する必要があり、NT形式(非輸送LHSおよび転置RHS)のみをサポートします。転置または他のFP8鋳造操作については、それらを独立して以前のカーネルに実装または融合してください。ライブラリはいくつかの単純なPytorchユーティリティ関数を提供しますが、これらはパフォーマンスが遅くなる可能性がありますが、私たちの主な焦点はGEMMカーネル自体の最適化にあります。

通常の濃い宝石(グループ化されていない)

基本的な非グループ化されたFP8 GEMMを実行するには、電話してください deep_gemm.gemm_fp8_fp8_bf16_nt 関数。詳細については、関数のドキュメントを参照してください。

グループ化されたgemms(連続レイアウト)

Cutlassの従来のグループ化されたGEMMとは異なり、DeepGEMMはm軸のみをグループ化しますが、NとKは固定されたままです。このデザインは、MOEモデルの専門家が同じ形状を共有するシナリオに合わせて調整されています。

各専門家がさまざまな数のトークンを処理する可能性のある前方パスまたは推論の予備をトレーニングするために、これらのトークンを「隣接する」レイアウトと呼ばれる単一のテンソルに連結します。各エキスパートセグメントは、GEMM Mブロックサイズに整列する必要があることに注意してください(get_m_alignment_for_contiguous_layout())。

詳細については、を参照してください m_grouped_gemm_fp8_fp8_bf16_nt_contiguous 関数ドキュメント。

グループ化されたgemms(マスクレイアウト)

推論デコードフェーズ中、CUDAグラフが有効になり、CPUが各専門家が受け取るトークンの数を知らない場合、マスクされたグループ化されたGEMMをサポートします。マスクテンソルを提供することにより、カーネルは有効な部分のみを計算します。

使用 m_grouped_gemm_fp8_fp8_bf16_nt_masked この目的のために、関連するドキュメントを参照してください。例の使用法は、からの低遅延カーネルの出力をからの出力を使用することです 深い 入力として。

ライブラリは、上記のカーネル以外にいくつかのユーティリティ関数を提供します。

  • deep_gemm.set_num_sms:使用する最大SMカウントを設定します
  • deep_gemm.get_num_sms:現在のSMの最大カウントを取得します
  • deep_gemm.get_m_alignment_for_contiguous_layout:グループ化された連続レイアウトのグループレベルのアライメント要件を取得する
  • deep_gemm.get_tma_aligned_size:必要なTMAアライメントサイズを取得します
  • deep_gemm.get_col_major_tma_aligned_tensor:カラム-MARJOR TMAに並べられたテンソルを取得します

ライブラリはまた、いくつかの環境変数を提供します。これは役立つ可能性があります。

  • DG_CACHE_DIR:文字列、コンパイルされたカーネルを保存するキャッシュディレクトリ、 $HOME/.deep_gemm デフォルトで
  • DG_NVCC_COMPILER:文字列、指定されたNVCCコンパイラパス。で見つかります from torch.utils.cpp_extension.CUDA_HOME デフォルトで
  • DG_DISABLE_FFMA_INTERLEAVE:0または1、FFMA-interLeavingの最適化を無効にします
  • DG_PTXAS_VERBOSE:0または1、詳細なptxasコンパイラ出力を表示します
  • DG_PRINT_REG_REUSE:0または1、FFMA-interLeavingの詳細を印刷します
  • DG_JIT_PRINT_NVCC_COMMAND:0または1、NVCCコンパイルコマンドを印刷します
  • DG_JIT_DEBUG:0または1、さらにデバッグ情報を印刷します

その他の例と詳細については、参照してください テストコード または、対応するPythonドキュメントを確認します。

cutlassから除外された技術を🐳で示します。

永続的なワープ特異化

Cutlassの設計に続いて、DeepGEMMのカーネルはゆがんだ特別なものであり、データの動き、テンソルコアMMA命令、およびCUDAコアプロモーションを重視しています。このプロセスを示す単純化された図を以下に示します。

テンソルメモリアクセラレータ (TMA)は、Hopper Architectureによって導入された新しいハードウェア機能で、より速く非同期データの動きのために設計されています。具体的には、TMAを使用します。

  • LHS、LHSスケーリング係数、およびRHSマトリックスのTMA負荷
  • 出力マトリックス用のTMAストア
  • TMAマルチキャスト(LHSマトリックス専用)
  • TMA記述子プリフェッチ

一般的な詳細最適化

統一された最適化されたブロックスケジューラ

DeepGEMMは完全に採用しています ジャストインタイム (JIT)設計、インストール時にコンピレーションは必要ありません。すべてのカーネルは、軽量のJIT実装を使用して実行時にコンパイルされます。このアプローチはいくつかの利点を提供します。

  • GEMMの形状、ブロックサイズ、およびパイプラインステージの数は、コンパイル時間定数として扱われます
    • レジスタの保存
    • コンパイラはより多くの最適化を行う場合があります
  • ブロックサイズ、ワープグループ数、最適なパイプラインステージ、およびTMAクラスターサイズの自動選択
    • しかし、自動調整がなければ、最適なものは決定論的に選択されます
  • MMAパイプラインを完全に展開し、コンパイラに最適化の機会を提供します
    • 小さな形状にとって非常に重要です
    • 参照してください launch_k_iterationsカーネルファイル 詳細については

全体として、JITは、のアプローチと同様に、小さな形状のパフォーマンスを大幅に向上させます トリトン コンパイラ。

アライメントされていないブロックサイズ🐳

特定の形状の場合、2のパワーに整列するブロックサイズは、十分に活用されていないSMSにつながる可能性があります。たとえば、 M=256, N=7168、の典型的なブロックサイズの割り当て BLOCK_M=128, BLOCK_N=128 結果のみになります (256 / 128) * (7168 / 128) = 112 利用されている132のSMSのうち。これに対処するために、112のような整理されていないブロックサイズをサポートし、有効にします (256 / 128) * (7168 / 112) = 128 このようなシナリオで動作するSMS。この手法を微調整されたスケーリングとともに実装するには、慎重に最適化する必要がありますが、最終的にはパフォーマンスの向上を実現します。

ffma sassインターリーブ🐳

パフォーマンスの改善が観察されます Cutlass FP8カーネル NVCC 12.2と12.3の間。コンパイルされたSASSを比較することで、 一連の FADD 説明書 インターリーブパターンで反転します。いくつかのオープンソースを参照した後 CUDAは組み立てられます 実装では、このビットが制御されることを特定しました yield、これは、ワープレベルの並列性を強化する可能性があります(推測だけで、現在のワープを生み出し、他の縦糸を機能させます)。

これを活用するために、私たちは開発します 同様のスクリプト を変更するには FFMA コンパイルされたバイナリの指示。単に変更するだけでなく yield ビット、私たちもひっくり返します reuse BIT(ワープが生成された場合、レジスタを再利用できません)。この調整により、MMAの指示をプロモーションと重複させる機会を増やすことにより、微調整されたスケーリングFP8 GEMMのパフォーマンス(場合によっては10%以上)が向上します FFMA 説明書。

deepgemmはに触発されています カトラス プロジェクト。開発者に感謝します!

このコードリポジトリは下にリリースされます MITライセンス

@misc{deepgemm2025,
      title={DeepGEMM: clean and efficient FP8 GEMM kernels with fine-grained scaling}, 
      author={Chenggang Zhao and Liang Zhao and Jiashi Li and Zhean Xu},
      year={2025},
      publisher = {GitHub},
      howpublished = {url{https://github.com/deepseek-ai/DeepGEMM}},
}

#deepseekaideepgemmdeepgemmきれいで効率的なfp8ジェムカーネル

執筆者について: nipponese

Nipponese News編集部は、国内外のニュースを日本語で分かりやすくお届けします。