Request for batched general_gemm() (or FP8-aware torch.bmm) for non-Linear GEMM workloads
Los mantenedores suelen responder en 2 días
Nadie ha tomado este issue todavía.
Evaluación
- Dificultad
- 5/5
- Tiempo estimado
- Más de una semana
- Aptitud para principiantes
- 35/100
- Tipo de issue
- Nueva funcionalidad
- Claridad
- Bastante claro
- Estado de actividad
- Tranquilo
- Área
- machine-learning, performance
Línea de trabajo
Comienza con general_gemm en transformer_engine.pytorch.cpp_extensions y compara su contrato 2D con las entradas 3D de torch.bmm. Revisa las rutas existentes de Float8Tensor y MXFP8Tensor, incluido el uso de use_split_accumulator y la compatibilidad con backward. Se considera terminado cuando los GEMMs FP8 por lotes con acumulación FP32 estén disponibles para la carga de trabajo de entrenamiento descrita sin iterar sobre slices del batch.
Escrito por el modelo de indexación a partir del texto del issue.
Descripción
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.
- Lenguaje dominante
- Python
- Estrellas
- 3.6k
- Forks
- 851
- Merge medio
- 4 d 15 h
- PR fusionados (30 d)
- 51
Preparar el entorno
- Sin Dockerfile ni archivo de Docker Compose
- Tiene una plantilla de pull request
- Leer la guía de contribución
Primeros pasos
- Lee el issue completo y luego la guía de contribución del proyecto.
- Comenta en el issue que vas a ocuparte — evita que dos personas hagan lo mismo.
- Haz un fork del repositorio y trabaja en una rama.
- Abre un pull request que haga referencia al número del issue.
Más de NVIDIA/TransformerEngine
-
[PyTorch] fp8_cs_quantize fake implementation returns a vector inverse scale instead of a scalarAbierto
Dificultad 2/5 1-3 horas Aptitud para principiantes 82/100
NVIDIA/TransformerEngine#3636 ·
Los mantenedores suelen responder en 2 días
-
[Bug] Backend selection picks FA3 for training with head_dim_qk=192 / v_head_dim=128, but FA3 backward cannot run itPosiblemente ocupada @yuweih205 la tomó hace 31 días. Abiertoattention
Dificultad 2/5 1-3 horas Aptitud para principiantes 85/100
NVIDIA/TransformerEngine#3481 · 4 comentarios ·
Los mantenedores suelen responder en 2 días
-
Increase MAX_TENSOR_NUMAbiertobug
Dificultad 2/5 1-3 horas Aptitud para principiantes 68/100
NVIDIA/TransformerEngine#2189 · 7 comentarios · 5 reacciones ·
Los mantenedores suelen responder en 2 días
-
Dificultad 4/5 3-5 días Aptitud para principiantes 50/100
NVIDIA/TransformerEngine#3645 ·
Los mantenedores suelen responder en 2 días
-
enhancement
Dificultad 5/5 Más de una semana Aptitud para principiantes 35/100
NVIDIA/TransformerEngine#3644 ·
Los mantenedores suelen responder en 2 días
Todos los issues de NVIDIA/TransformerEngine
Issues similares
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 85/100
Los mantenedores suelen responder en 1 día
-
SR_SECURITY_DESCRIPTOR.fromString drops the SACL when no DACL is presentPosiblemente ocupada @paul7436 la tomó hoy. Abierto
Dificultad 2/5 1-3 horas Aptitud para principiantes 88/100
Los mantenedores suelen responder en 2 días
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 85/100
equinor/fmu-sumo-uploader#302 ·
Los mantenedores suelen responder en 1 día
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 75/100
modelscope/evalscope#1821 ·
Los mantenedores suelen responder en 1 día
-
Sanity on ansible-core devel fails: ignore-2.23.txt references the removed import-3.9 testPosiblemente ocupada @yurnov la tomó hoy. Abiertoneeds_triage
Dificultad 1/5 Menos de una hora Aptitud para principiantes 91/100
ansible-collections/kubernetes.core#1275 ·
Los mantenedores suelen responder en 1 día