FP8 DPA future token leakage
I maintainer di solito rispondono entro 2 giorni
Nessuno ha ancora preso questa issue.
Valutazione
- Difficoltà
- 5/5
- Tempo stimato
- Più di una settimana
- Idoneità per principianti
- 38/100
- Tipo di issue
- Bug
- Chiarezza
- Abbastanza chiara
- Stato di attività
- Tranquilla
- Ambito
- backend, machine-learning, performance, testing-qa
Direzione di ricerca
Esegui prima la riproduzione minima, quindi esamina i percorsi di attention citati in pytorch/attention/dot_product_attention/utils.py, common/fused_attn/fused_attn_fp8.cu e nei file di scaling delle recipe. Confronta il comportamento di MXFP8 e dello scaling corrente con i controlli BF16; il lavoro è completato quando gli output del prefisso causale rimangono invariati modificando solo i valori futuri, con copertura di regressione per entrambi i percorsi.
Scritto dal modello di indicizzazione a partire dal testo della issue.
Descrizione
FP8 scaling can leak future tokens into causal SDPA outputs
Summary
Transformer Engine has two current-data FP8 DPA paths in which scale selection can inspect masked future tokens before the causal mask is applied. MXFP8 quantizes V for PV with an E8M0 scale shared by 32 sequence positions, while tensorwise current scaling quantizes the combined Q/K/V tensor with one current amax. In both cases, a future V value can alter the quantized representation of an allowed prefix V value.
Expected behavior
For causal self-attention with identical Q, K, and V through position t, changing V only at positions greater than t should not change outputs through t:
O(Q, K, V_a)[:, :, :t+1, :] == O(Q, K, V_b)[:, :, :t+1, :]
when V_a[:, :, :t+1, :] == V_b[:, :, :t+1, :]
Small nondeterministic kernel variation aside, future masked values should not be an input to the numerical representation of allowed prefix values.
Actual behavior and causes
MXFP8 sequence-block V scaling
For each V block B_b = {32b, ..., 32b+31} and each batch, head, and V feature, TE computes:
D[b,h,d] = E8M0_round_up(max(abs(V[j,h,d]), j in B_b) / FP8_MAX)
V8[j,h,d] = cast_fp8(V[j,h,d] / D[b,h,d])
V_hat[j,h,d] = fp32(V8[j,h,d]) * D[b,h,d]
Causal PV then computes:
O[t,h,d] = sum(P[t,h,j] * V_hat[j,h,d], j <= t)
P is not the source of this dependency. In the MXFP8 path, TE nulls the external S and dP quantizer slots, and cuDNN's checked-in online-softmax reference quantizes the unnormalized P tile using fixed, data-independent constants s_scale=16 and inv_s_scale=1/16. The sequence-dependent reduction at issue is V's columnwise E8M0 amax.
A future V[r,h,d], where t < r < 32b+32, can change D[b,h,d], which can change the quantized representation of an allowed V[j,h,d] for j <= t. The future token has no direct PV term because P[t,h,r] = 0, but O[t,h,d] can still change through the shared scale.
Tensorwise current QKV scaling
The current-scaling path has the same causal problem with broader scope. For ordinary dense DPA, TE combines Q, K, and V and applies one current quantizer to the combined tensor:
A = max(abs(concat(Q, K, V)))
s = 448 / A
Q8,K8,V8 = round_E4M3(s * concat(Q, K, V))
Q_hat,... = fp32(Q8,K8,V8) / s
Therefore any future Q, K, or V outlier can change the representation of every prefix Q, K, and V value in the same local DPA tensor.
Source evidence
- TE selects rowwise Q/K but columnwise V for forward MXFP8 attention, so V scales reduce along S:
utils.pylines 2693–2727. - TE passes V's columnwise payload and columnwise scale factors to the fused backend:
fused_attn_fp8.culines 1105–1123. - cuDNN Frontend fixes the block size at 32 and validates V scales as
[b, h, ceil(s_kv/32), d]:scaled_dot_product_flash_attention.hlines 326–413. - cuDNN dequantizes V with block size
{1, 32}immediately before BMM2/PV:scaled_dot_product_flash_attention.hlines 1026–1046. - The causal mask is applied to the QK score tensor inside SDPA, after V has already been quantized externally:
scaled_dot_product_flash_attention.hlines 860–884. - For non-MX FP8 DPA, TE combines packed or separate Q, K, and V storage and invokes the QKV quantizer once on the combined tensor:
utils.pylines 2777–2813. - Current scaling computes
max_fp8 / amaxas a full FP32 value by default and clears the mantissa only when power-of-two scaling is explicitly enabled:recipe_common.cuhlines 14–48 andFloat8CurrentScalingdefaults at lines 284–304.
Minimal repro
import os
os.environ["NVTE_ALLOW_NONDETERMINISTIC_ALGO"] = "0"
os.environ["NVTE_FLASH_ATTN"] = "0"
os.environ["NVTE_FUSED_ATTN"] = "1"
os.environ["NVTE_UNFUSED_ATTN"] = "0"
import torch
import transformer_engine
from transformer_engine.common.recipe import Float8CurrentScaling, Format, MXFP8BlockScaling
from transformer_engine.pytorch import DotProductAttention, autocast
S, D, T = 128, 128, 7
q = torch.ones((S, 1, 1, D), device="cuda", dtype=torch.bfloat16)
k = torch.ones_like(q)
mx_base_v = torch.zeros_like(q)
mx_base_v[0, 0, 0, 0] = 2.0**-20
mx_same_block_v = mx_base_v.clone()
mx_same_block_v[16, 0, 0, 0] = 1.0
mx_next_block_v = mx_base_v.clone()
mx_next_block_v[32, 0, 0, 0] = 1.0
current_base_v = torch.zeros_like(q)
current_base_v[0, 0, 0, 0] = 2.0**-10
current_future_v = current_base_v.clone()
current_future_v[64, 0, 0, 0] = 1024.0
mx_recipe = MXFP8BlockScaling(
fp8_format=Format.E4M3,
fp8_dpa=True,
fp8_mha=False,
)
current_recipe = Float8CurrentScaling(
fp8_format=Format.E4M3,
fp8_dpa=True,
fp8_mha=False,
)
def run(value, fp8, recipe=None):
dpa = DotProductAttention(
1,
D,
attention_dropout=0.0,
qkv_format="sbhd",
attn_mask_type="causal",
).cuda().eval()
with torch.no_grad(), autocast(enabled=fp8, recipe=recipe):
return dpa(q, k, value, attn_mask_type="causal").float()
def report(name, reference, changed):
delta = (reference[: T + 1] - changed[: T + 1]).abs()
print(
f"{name:22} max_abs={delta.max().item():.10e} "
f"changed={delta.count_nonzero().item():2d} "
f"O[t,0]={reference[T, 0, 0].item():.10e} -> "
f"{changed[T, 0, 0].item():.10e}"
)
return delta.max().item()
mx_base = run(mx_base_v, True, mx_recipe)
mx_same = run(mx_same_block_v, True, mx_recipe)
mx_next = run(mx_next_block_v, True, mx_recipe)
current_base = run(current_base_v, True, current_recipe)
current_future = run(current_future_v, True, current_recipe)
bf16_mx_base = run(mx_base_v, False)
bf16_mx_same = run(mx_same_block_v, False)
bf16_current_base = run(current_base_v, False)
bf16_current_future = run(current_future_v, False)
print(
f"TE={transformer_engine.__version__} torch={torch.__version__} "
f"cuDNN={torch.backends.cudnn.version()} GPU={torch.cuda.get_device_name(0)}"
)
same_block_diff = report("MXFP8 same block", mx_base, mx_same)
next_block_diff = report("MXFP8 next block", mx_base, mx_next)
current_diff = report("Current future token", current_base, current_future)
bf16_mx_diff = report("BF16 MX control", bf16_mx_base, bf16_mx_same)
bf16_current_diff = report(
"BF16 current control", bf16_current_base, bf16_current_future
)
assert next_block_diff == 0.0
assert bf16_mx_diff == 0.0
assert bf16_current_diff == 0.0
failures = []
if same_block_diff != 0.0:
failures.append("MXFP8 same-block V scale changed the causal prefix")
if current_diff != 0.0:
failures.append("current QKV tensor scale changed the causal prefix")
assert not failures, "; ".join(failures)
Observed current behavior
Reproduced on B200, Transformer Engine 2.16.0+4220403e, PyTorch 2.13.0a0+8145d630e8.nv26.06, and cuDNN 9.23.0:
TE=2.16.0+4220403e torch=2.13.0a0+8145d630e8.nv26.06 cuDNN=92300 GPU=NVIDIA B200
MXFP8 same block max_abs=9.5367431641e-07 changed= 8 O[t,0]=1.1920928955e-07 -> 0.0000000000e+00
MXFP8 next block max_abs=0.0000000000e+00 changed= 0 O[t,0]=1.1920928955e-07 -> 1.1920928955e-07
Current future token max_abs=9.7656250000e-04 changed= 8 O[t,0]=1.2207031250e-04 -> 0.0000000000e+00
BF16 MX control max_abs=0.0000000000e+00 changed= 0 O[t,0]=1.1920928955e-07 -> 1.1920928955e-07
BF16 current control max_abs=0.0000000000e+00 changed= 0 O[t,0]=1.2207031250e-04 -> 1.2207031250e-04
AssertionError: MXFP8 same-block V scale changed the causal prefix; current QKV tensor scale changed the causal prefix
- Lingua principale
- Python
- Stelle
- 3.6k
- Fork
- 851
- Merge medio
- 4g 15h
- PR unite (30g)
- 51
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
-
[PyTorch] fp8_cs_quantize fake implementation returns a vector inverse scale instead of a scalarAperta
Difficoltà 2/5 1-3 ore Idoneità per principianti 82/100
NVIDIA/TransformerEngine#3636 ·
I maintainer di solito rispondono entro 2 giorni
-
[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 31 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
-
Difficoltà 4/5 3-5 giorni Idoneità per principianti 50/100
NVIDIA/TransformerEngine#3645 ·
I maintainer di solito rispondono entro 2 giorni
-
enhancement
Difficoltà 5/5 Più di una settimana Idoneità per principianti 35/100
NVIDIA/TransformerEngine#3644 ·
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 72/100
Graphify-Labs/graphify#4241 · 1 commento ·
I maintainer di solito rispondono entro 1 giorno
-
Difficoltà 1/5 Meno di un'ora Idoneità per principianti 72/100
-
DeviceTrackerAperta
Difficoltà 2/5 1-3 ore Idoneità per principianti 63/100
XiaoMi/ha_xiaomi_home#1821 ·
I maintainer di solito rispondono entro 1 giorno
-
Maven path-index: "Ambiguous or noncanonical artifact path" error does not report the offending pathAperta
Difficoltà 2/5 1-3 ore Idoneità per principianti 76/100
pulp/pulp_maven#524 ·
I maintainer di solito rispondono entro 1 giorno
-
Difficoltà 1/5 1-3 ore Idoneità per principianti 82/100
I maintainer di solito rispondono entro 1 giorno