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は歓迎されます。
| 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 |
| #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 |
| #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 関数。詳細については、関数のドキュメントを参照してください。
Cutlassの従来のグループ化されたGEMMとは異なり、DeepGEMMはm軸のみをグループ化しますが、NとKは固定されたままです。このデザインは、MOEモデルの専門家が同じ形状を共有するシナリオに合わせて調整されています。
各専門家がさまざまな数のトークンを処理する可能性のある前方パスまたは推論の予備をトレーニングするために、これらのトークンを「隣接する」レイアウトと呼ばれる単一のテンソルに連結します。各エキスパートセグメントは、GEMM Mブロックサイズに整列する必要があることに注意してください(get_m_alignment_for_contiguous_layout())。
詳細については、を参照してください m_grouped_gemm_fp8_fp8_bf16_nt_contiguous 関数ドキュメント。
推論デコードフェーズ中、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記述子プリフェッチ
- の使用
stmatrixPTX命令 - カウントコントロールを登録します さまざまなワープグループに合わせて調整されています
- 可能な限り重複する、例えばTMAストアと非TMA RHSスケーリング因子負荷🐳
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。この手法を微調整されたスケーリングとともに実装するには、慎重に最適化する必要がありますが、最終的にはパフォーマンスの向上を実現します。
パフォーマンスの改善が観察されます 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ジェムカーネル