FP8 DPA future token leakage
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
- 38/100
- Tipo de issue
- Error
- Claridad
- Bastante claro
- Estado de actividad
- Tranquilo
- Área
- backend, machine-learning, performance, testing-qa
Línea de trabajo
Ejecuta primero la reproducción mínima y luego inspecciona las rutas de atención citadas en pytorch/attention/dot_product_attention/utils.py, common/fused_attn/fused_attn_fp8.cu y los archivos de escalado de recipe. Compara el comportamiento de MXFP8 y del escalado actual con los controles BF16; se considera terminado cuando las salidas del prefijo causal permanecen sin cambios al modificar únicamente los valores futuros, con cobertura de regresión para ambas rutas.
Escrito por el modelo de indexación a partir del texto del issue.
Descripción
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
- Lenguaje dominante
- Python
- Estrellas
- 3.6k
- Forks
- 851
- Merge medio
- 5 d 1 h
- PR fusionados (30 d)
- 52
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
-
[Bug] group_quantize_fp8_blockwise: mbarrier invalidated before other threads finish waiting on itAbierto
Dificultad 2/5 1-3 horas Aptitud para principiantes 82/100
NVIDIA/TransformerEngine#3647 ·
Los mantenedores suelen responder en 2 días
-
[PyTorch] fp8_cs_quantize fake implementation returns a vector inverse scale instead of a scalarPosiblemente ocupada @sanjana658 la tomó hoy. Abierto
Dificultad 2/5 1-3 horas Aptitud para principiantes 82/100
NVIDIA/TransformerEngine#3636 · 2 comentarios ·
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
-
[PyTorch] CUDA graph RNG registration floods training logs on automatic-registration buildsPosiblemente ocupada @ksivaman la tomó hoy. Abierto
Dificultad 4/5 3-5 días Aptitud para principiantes 50/100
NVIDIA/TransformerEngine#3645 · 1 comentario · 1 asignado ·
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 68/100
pyjanitor-devs/pyjanitor#1758 ·
Los mantenedores suelen responder en 1 día
-
bug ready for review
Dificultad 2/5 1-3 horas Aptitud para principiantes 86/100
odysseus-dev/odysseus#6641 ·
Los mantenedores suelen responder en 1 día
-
bug
Dificultad 2/5 1-3 horas Aptitud para principiantes 76/100
happypawspillaro/happypaws#78 ·
Los mantenedores suelen responder en 4 días
-
pydanty:is-working
Dificultad 2/5 1-3 horas Aptitud para principiantes 82/100
pydantic/pydantic-ai#10020 ·
Los mantenedores suelen responder en 1 día
-
stdlib type-bug
Dificultad 2/5 1-3 horas Aptitud para principiantes 68/100
python/cpython#159044 · 4 comentarios ·
Los mantenedores suelen responder en 1 día