Request for batched general_gemm() (or FP8-aware torch.bmm) for non-Linear GEMM workloads
メンテナーはふだん 2 日以内に返信
まだ誰も着手していません。
評価
- 難易度
- 5/5
- 見積もり時間
- 1週間以上
- 初心者へのやさしさ
- 35/100
- issue の種類
- 機能追加
- 明瞭さ
- おおむね明確
- 活発さ
- 静か
調査の方向性
transformer_engine.pytorch.cpp_extensions の general_gemm から始め、その 2D コントラクトを torch.bmm の 3D 入力と比較します。既存の Float8Tensor および MXFP8Tensor のパスを確認し、use_split_accumulator の使用と backward のサポートも含めます。バッチスライスをループせずに、説明されているトレーニングワークロードで FP32 アキュムレーションを使用するバッチ化 FP8 GEMM が利用可能になれば完了です。
索引モデルが issue の本文から書いたものです。
説明
Is your feature request related to a problem? Please describe.
We’re accelerating triangular multiplication in a protein structure prediction model (AlphaFold-style tri-mul). The core operation is two large einsums over 4D pair representations that we’ve reshaped into batched matmuls:
# Input: (B, N, N, D) where N = 2048 (sequence length), D = 128
# After chunk, permute, reshape: (B*32, 2048, 2048)
x1 = torch.bmm(a, b.transpose(1, 2)) # B*32 independent N×N GEMMs
At N = 2048, this accounts for roughly 40% of the tri-mul compute and is heavily memory-bandwidth-bound. Currently we run in FP32 (4 bytes/element) or BF16 (2 bytes/element). MXFP8 inputs (1 byte/element) with FP32 accumulation would provide up to a 4× reduction in HBM reads, which is the dominant cost at these sizes.
However, there is currently no way to run FP8 batched matrix multiplication through TE:
te.autocast()only intercepts TE modules, nottorch.bmmFloat8Tensor/MXFP8Tensorpassed totorch.bmmsilently dequantize to full precisiongeneral_gemm()supports FP8 × FP8 withuse_split_accumulator=True, but only accepts 2D inputs — looping overB*32slices would likely negate the bandwidth savings
Related: #1910 describes the same gap for FP8 GEMM beyond te.Linear.
Describe the solution you’d like
A batched variant of general_gemm() that accepts 3D inputs and runs FP8 GEMMs across the batch dimension with FP32 accumulation:
from transformer_engine.pytorch.cpp_extensions import batched_general_gemm
# Quantize inputs to FP8
a_fp8 = mxfp8_quantizer.quantize(a_3d) # (B*32, N, N)
b_fp8 = mxfp8_quantizer.quantize(b_3d) # (B*32, N, N)
# Batched FP8 GEMM with FP32 accumulation
output = batched_general_gemm(
a_fp8,
b_fp8,
out_dtype=torch.bfloat16,
layout="NN",
use_split_accumulator=True, # FP8×FP8 multiply, FP32 accumulate
)
# output: (B*32, N, N) in BF16
Alternatively, making Float8Tensor / MXFP8Tensor dispatch torch.bmm to real FP8 tensor core GEMMs, instead of dequantizing, would also solve this.
Describe alternatives you’ve considered
- GroupedLinear: Suggested in
#1910, but it is designed for MoE-style use cases with different weights per group. Our use case is two arbitrary input tensors, not input × stored weight. It was also noted there may be significant overhead. - Looping
general_gemm()over batch slices: Functionally possible, but Python loop overhead and the lack of kernel batching would likely wipe out the memory-bandwidth gains from FP8. - Skipping
.float()and runningtorch.bmmin BF16: This is our current workaround. It gives a 2× memory reduction versus FP32, but still leaves another 2× on the table compared with FP8.
Additional context
- Targeting Blackwell (
MXFP8BlockScaling) and Hopper (DelayedScaling/CurrentScaling) - Training workload, so backward-pass support is needed
- This batched FP8 GEMM pattern would also help other workloads with non-
Linearmatmuls, including attention (unfused path), structure prediction, graph neural networks, and any model with einsum contractions reshaped tobmm - TE v1.12+
Happy to provide a minimal repro or benchmark if helpful.
- 主要言語
- Python
- スター
- 3.6k
- フォーク
- 851
- 平均マージ
- 4日 15時間
- マージ済み PR(30日)
- 51
環境構築
- Dockerfile・Docker Compose ファイルなし
- プルリクエストのテンプレートあり
- コントリビューションガイドを読む
はじめの一歩
- issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
- 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
- リポジトリをフォークし、ブランチを切って変更します。
- issue 番号を参照したプルリクエストを送ります。
NVIDIA/TransformerEngine のほかの issue
-
[PyTorch] fp8_cs_quantize fake implementation returns a vector inverse scale instead of a scalarオープン
難易度 2/5 1〜3時間 初心者へのやさしさ 82/100
NVIDIA/TransformerEngine#3636 ·
メンテナーはふだん 2 日以内に返信
-
[Bug] Backend selection picks FA3 for training with head_dim_qk=192 / v_head_dim=128, but FA3 backward cannot run it対応中かも @yuweih205 が 31 日前に担当しました。 オープンattention
難易度 2/5 1〜3時間 初心者へのやさしさ 85/100
NVIDIA/TransformerEngine#3481 · コメント 4 件 ·
メンテナーはふだん 2 日以内に返信
-
bug
難易度 2/5 1〜3時間 初心者へのやさしさ 68/100
NVIDIA/TransformerEngine#2189 · コメント 7 件 · リアクション 5 件 ·
メンテナーはふだん 2 日以内に返信
-
難易度 4/5 3〜5日 初心者へのやさしさ 50/100
NVIDIA/TransformerEngine#3645 ·
メンテナーはふだん 2 日以内に返信
-
enhancement
難易度 5/5 1週間以上 初心者へのやさしさ 35/100
NVIDIA/TransformerEngine#3644 ·
メンテナーはふだん 2 日以内に返信
NVIDIA/TransformerEngine の issue をすべて見る
似ている issue
-
難易度 2/5 1〜3時間 初心者へのやさしさ 82/100
RedHatQE/mtv-api-tests#721 ·
メンテナーはふだん 1 日以内に返信
-
難易度 2/5 1〜3時間 初心者へのやさしさ 84/100
メンテナーはふだん 1 日以内に返信
-
難易度 1/5 1〜3時間 初心者へのやさしさ 85/100
pytest-dev/pluggy#757 ·
メンテナーはふだん 1 日以内に返信
-
難易度 1/5 1〜3時間 初心者へのやさしさ 85/100
NousResearch/hermes-agent#134960 ·
メンテナーはふだん 1 日以内に返信
-
HTML backend: `<br>` leaks the internal sentinel U+E000 into list items, headings and captions対応中かも @morten-lagabote が今日担当しました。 オープン
難易度 2/5 1〜3時間 初心者へのやさしさ 67/100
docling-project/docling#4671 ·
メンテナーはふだん 1 日以内に返信