[bug] NVFP4 + `torch.compile`: errors with 3D input
I maintainer di solito rispondono entro 2 giorni
Valutazione
- Difficoltà
- 2/5
- Tempo stimato
- 1-3 ore
- Idoneità per principianti
- 25/100
- Tipo di issue
- Bug
- Chiarezza
- Specificata chiaramente
- Stato di attività
- Ferma
- Ambito
- machine-learning
Direzione di ricerca
Start with transformer_engine/pytorch/tensor/nvfp4_tensor.py, especially NVFP4Quantizer.get_columnwise_shape(), and compare its returned shape with the columnwise buffer shape described in the issue. The issue reports forward and backward verification on GB200 for 3D and 2D NVFP4 inputs and 3D MXFP8 input; done means those cases pass without regression under torch.compile. A linked pull request is already open.
Scritto dal modello di indicizzazione a partire dal testo della issue.
Descrizione
Summary
Under torch.compile, an NVFP4-quantized te.pytorch.Linear fails with
AssertionError: wrong number of dimensions2 for op:
torch.ops.transformer_engine_compile.linear.default
whenever the input activation has 3 (e.g. the normal
(batch, seq, hidden) shape of any real transformer forward pass). A 2D
input (tokens, hidden) does not trigger it, which is why this can look
like it works in isolated/synthetic benchmarks and only surfaces once a
real model is compiled end-to-end.
Environment
- TransformerEngine:
2.21.0.dev0+9233c73 - PyTorch:
2.15.0a0+git8c12b70 - GPU: GB200
- Recipe:
transformer_engine.common.recipe.NVFP4BlockScaling()with RHT
enabled (the default)
Minimal repro
import torch
import transformer_engine.pytorch as te
from transformer_engine.common import recipe
hidden = 5120
model = te.Linear(hidden, hidden, params_dtype=torch.bfloat16, device="cuda")
x = torch.randn(1, 4096, hidden, dtype=torch.bfloat16, device="cuda", requires_grad=True)
def step(inp):
with te.autocast(recipe=recipe.NVFP4BlockScaling()):
return model(inp).sum()
compiled_step = torch.compile(step, fullgraph=False)
# loss = step(x) # works
loss = compiled_step(x)
loss.backward()
Observed failure
AssertionError: wrong number of dimensions2 for op:
torch.ops.transformer_engine_compile.linear.default
assert_tensor_metadata(buf17, (5120, 1, 2048), (2048, 2048, 1), torch.uint8,
'torch.ops.transformer_engine_compile.linear.default')
Root cause (AI generated)
NVFP4Quantizer.get_columnwise_shape()
(transformer_engine/pytorch/tensor/nvfp4_tensor.py:309-325) computes the
shape of the RHT/columnwise-quantized buffer as:
colwise_shape = [shape[-1]]
for i in range(len(shape) - 1):
colwise_shape.append(shape[i])
return tuple(colwise_shape)
i.e. it preserves the input's full rank, permuting each leading dimension
individually ((b, s, h) → (h, b, s)).
torch.compile's Inductor backend uses this function's result (via
inner_tensor_specs) as the compile-time-authoritative shape/stride for
the columnwise buffer, and bakes an assert_tensor_metadata /
assert_size_stride guard into the generated code against it.
But the real CUDA quantization kernel does not preserve rank: it always
collapses every leading dimension into a single M dimension and treats
the tensor as 2D (M, K), regardless of the logical input rank.
For a 2D input, there's only one leading dimension, so the two views
happen to coincide and the bug stays masked. For any ≥3D input — i.e. a
normal batched activation — the predicted (compile-time) shape and the
actual (runtime) shape diverge, and the Inductor-generated guard fails.
Proposed fix (AI generated)
Collapse the leading dimensions into one, matching what the real kernel
does:
--- a/transformer_engine/pytorch/tensor/nvfp4_tensor.py
+++ b/transformer_engine/pytorch/tensor/nvfp4_tensor.py
@@ -327,10 +327,19 @@ class NVFP4Quantizer(Quantizer):
if len(shape) == 0:
return tuple()
# and then after AG, a reorganize kernel will be called to restore the shape
- colwise_shape = [shape[-1]]
- for i in range(len(shape) - 1):
- colwise_shape.append(shape[i])
- return tuple(colwise_shape)
+ # Collapse all leading dims into one -- the real RHT/columnwise CUDA
+ # kernel always treats the input as a 2D (M, K) matrix (M = prod of
+ # leading dims), never preserving the original tensor rank. Returning
+ # a full-rank shape here (one entry per leading dim) made this
+ # function disagree with the real kernel output whenever the input
+ # had more than one leading dim (e.g. a (batch, seq, hidden)
+ # activation): torch.compile's Inductor backend registers this
+ # function's result as the expected compile-time shape/stride for
+ # the columnwise buffer, so the mismatch surfaced as
+ # "AssertionError: wrong number of dimensions2 for op:
+ # torch.ops.transformer_engine_compile.linear.default" the first time
+ # a >= 3D activation reached an NVFP4 te.Linear under torch.compile.
+ return (shape[-1], math.prod(shape[:-1]))
Verification of the fix
Verified on a real GB200 node, forward + backward, across:
- The originally-failing case: NVFP4, 3D activation,
torch.compile
(fullgraph=False, default mode) — now passes - NVFP4, 3D activation,
reduce-overheadmode — passes - NVFP4, 2D activation (pre-existing passing case) — still passes, no
regression - MXFP8, 3D activation (pre-existing passing case, different
quantizer) — still passes, no regression
No other code path needed changes — an initial backward-pass error seen
mid-investigation (RuntimeError: mismatch in length of strides and shape) turned out to be a stale torch.compile/Inductor disk cache
(/tmp/torchinductor_*) left over from a pre-patch run, not a second bug;
clearing the cache resolved it with only the one fix above.
cc: @pggPL
- Lingua principale
- Python
- Stelle
- 3.6k
- Fork
- 844
- Merge medio
- 4g 42m
- PR unite (30g)
- 49
Preparare l'ambiente
- Nessun Dockerfile né file Docker Compose
- Ha un modello di pull request
- Leggi la guida per i contributori
Come iniziare
- Leggi tutta la issue e poi la guida ai contributi del progetto.
- Commenta sulla issue per dire che te ne occupi tu — evita che due persone facciano lo stesso lavoro.
- Fai un fork del repository e lavora su un branch.
- Apri una pull request che faccia riferimento al numero della issue.
Altre issue di NVIDIA/TransformerEngine
-
[Bug] Backend selection picks FA3 for training with head_dim_qk=192 / v_head_dim=128, but FA3 backward cannot run itForse già presa @yuweih205 l’ha presa 29 giorni fa. Apertaattention
Difficoltà 2/5 1-3 ore Idoneità per principianti 85/100
NVIDIA/TransformerEngine#3481 · 4 commenti ·
I maintainer di solito rispondono entro 2 giorni
-
Increase MAX_TENSOR_NUMApertabug
Difficoltà 2/5 1-3 ore Idoneità per principianti 68/100
NVIDIA/TransformerEngine#2189 · 7 commenti · 5 reazioni ·
I maintainer di solito rispondono entro 2 giorni
-
[BUG] Grouped MXFP8 quantization is not concurrency safe with multiple streamsForse già presa @kainzhong l’ha presa oggi. Apertabug
Difficoltà 4/5 3-5 giorni Idoneità per principianti 25/100
NVIDIA/TransformerEngine#3630 ·
I maintainer di solito rispondono entro 2 giorni
-
Multi-tensor swizzle kernels fail with "too many resources requested for launch" (missing __launch_bounds__)Forse già presa @ravimajeti l’ha presa 2 giorni fa. Aperta
Difficoltà 2/5 1-3 ore Idoneità per principianti 25/100
NVIDIA/TransformerEngine#3621 ·
I maintainer di solito rispondono entro 2 giorni
-
[PyTorch] Avoid selecting FA4 for deterministic training on SM120Forse già presa Una pull request collegata a questa issue è aperta o già unita. Apertabug
Difficoltà 3/5 1-2 giorni Idoneità per principianti 35/100
NVIDIA/TransformerEngine#3594 ·
I maintainer di solito rispondono entro 2 giorni
Tutte le issue di NVIDIA/TransformerEngine
Issue simili
-
Difficoltà 2/5 1-3 ore Idoneità per principianti 83/100
I maintainer di solito rispondono entro 1 giorno
-
Difficoltà 2/5 1-3 ore Idoneità per principianti 86/100
FuRongJun-1999/dsh-memory#65 ·
I maintainer di solito rispondono entro 1 giorno
-
ci needs-ac
Difficoltà 2/5 1-3 ore Idoneità per principianti 75/100
Ikalus1988/MisakaNet#2930 ·
I maintainer di solito rispondono entro 1 giorno
-
`FakeBackendV2.run` fails with `NoiseError` on circuits with delays on qubits where T2 > 2·T1Apertabug
Difficoltà 2/5 1-3 ore Idoneità per principianti 78/100
Qiskit/qiskit-aer#2466 ·
-
area/cli
Difficoltà 2/5 1-3 ore Idoneità per principianti 82/100