Hacktoberfest 2026:メンテナが10月に向けて印を付けた、オープンで初心者向けの issue。 Hacktoberfest の issue を見る

Request for batched general_gemm() (or FP8-aware torch.bmm) for non-Linear GEMM workloads

オープン
#2,846 コメント 0 件 リアクション 0 件 担当者 0 名 GitHub で見る

メンテナーはふだん 2 日以内に返信

まだ誰も着手していません。

評価

難易度
5/5
見積もり時間
1週間以上
初心者へのやさしさ
35/100
issue の種類
機能追加
明瞭さ
おおむね明確
活発さ
静か
技術スタック
python, pytorch

調査の方向性

transformer_engine.pytorch.cpp_extensions の general_gemm から始め、その 2D コントラクトを torch.bmm の 3D 入力と比較します。既存の Float8Tensor および MXFP8Tensor のパスを確認し、use_split_accumulator の使用と backward のサポートも含めます。バッチスライスをループせずに、説明されているトレーニングワークロードで FP32 アキュムレーションを使用するバッチ化 FP8 GEMM が利用可能になれば完了です。

索引モデルが issue の本文から書いたものです。

説明

field-request

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, not torch.bmm
  • Float8Tensor / MXFP8Tensor passed to torch.bmm silently dequantize to full precision
  • general_gemm() supports FP8 × FP8 with use_split_accumulator=True, but only accepts 2D inputs — looping over B*32 slices 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 running torch.bmm in 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-Linear matmuls, including attention (unfused path), structure prediction, graph neural networks, and any model with einsum contractions reshaped to bmm
  • TE v1.12+

Happy to provide a minimal repro or benchmark if helpful.

主要言語
Python
スター
3.6k
フォーク
851
平均マージ
4日 15時間
マージ済み PR(30日)
51

環境構築

はじめの一歩

  1. issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
  2. 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
  3. リポジトリをフォークし、ブランチを切って変更します。
  4. issue 番号を参照したプルリクエストを送ります。

NVIDIA/TransformerEngine のほかの issue

NVIDIA/TransformerEngine の issue をすべて見る

似ている issue

Python の issue をもっと見る

新しい issue をメールで受け取る

初心者向けの GitHub issue を短くまとめたダイジェスト。