FP8 DPA future token leakage
Maintainer thường phản hồi trong vòng 2 ngày
Chưa có ai nhận issue này.
Đánh giá
- Độ khó
- 5/5
- Thời gian dự kiến
- Hơn một tuần
- Mức phù hợp với người mới
- 38/100
- Loại issue
- Lỗi
- Độ rõ ràng
- Khá rõ ràng
- Mức độ hoạt động
- Ít trao đổi
- Lĩnh vực
- backend, machine-learning, performance, testing-qa
Hướng nghiên cứu
Trước tiên, hãy chạy bản tái hiện tối thiểu, sau đó kiểm tra các đường dẫn attention được nêu trong pytorch/attention/dot_product_attention/utils.py, common/fused_attn/fused_attn_fp8.cu và các tệp scaling của recipe. So sánh hành vi của MXFP8 và scaling hiện tại với các kiểm soát BF16; được xem là hoàn tất khi đầu ra của tiền tố nhân quả không thay đổi khi chỉ các giá trị tương lai thay đổi, với độ bao phủ hồi quy cho cả hai đường dẫn.
Do mô hình lập chỉ mục viết ra từ nội dung của issue.
Mô tả
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
- Ngôn ngữ chính
- Python
- Star
- 3.6k
- Fork
- 851
- Merge trung bình
- 5 ngày 1 giờ
- Pull request đã merge (30 ngày)
- 52
Chuẩn bị môi trường
- Không có Dockerfile hay tệp Docker Compose
- Có mẫu pull request
- Đọc hướng dẫn đóng góp
Bắt đầu từ đâu
- Đọc hết issue, rồi đọc hướng dẫn đóng góp của dự án.
- Bình luận trên issue rằng bạn sẽ nhận — tránh hai người làm cùng một việc.
- Fork repository và làm thay đổi trên một nhánh.
- Mở pull request có tham chiếu số hiệu của issue.
Issue khác của NVIDIA/TransformerEngine
-
[Bug] group_quantize_fp8_blockwise: mbarrier invalidated before other threads finish waiting on itĐang mở
Độ khó 2/5 1-3 giờ Mức phù hợp với người mới 82/100
NVIDIA/TransformerEngine#3647 ·
Maintainer thường phản hồi trong vòng 2 ngày
-
[PyTorch] fp8_cs_quantize fake implementation returns a vector inverse scale instead of a scalarCó thể đã có người làm @sanjana658 đã nhận 1 ngày trước. Đang mở
Độ khó 2/5 1-3 giờ Mức phù hợp với người mới 82/100
NVIDIA/TransformerEngine#3636 · 2 bình luận ·
Maintainer thường phản hồi trong vòng 2 ngày
-
[Bug] Backend selection picks FA3 for training with head_dim_qk=192 / v_head_dim=128, but FA3 backward cannot run itCó thể đã có người làm @yuweih205 đã nhận 32 ngày trước. Đang mởattention
Độ khó 2/5 1-3 giờ Mức phù hợp với người mới 85/100
NVIDIA/TransformerEngine#3481 · 4 bình luận ·
Maintainer thường phản hồi trong vòng 2 ngày
-
Increase MAX_TENSOR_NUMĐang mởbug
Độ khó 2/5 1-3 giờ Mức phù hợp với người mới 68/100
NVIDIA/TransformerEngine#2189 · 7 bình luận · 5 reaction ·
Maintainer thường phản hồi trong vòng 2 ngày
-
[PyTorch] CUDA graph RNG registration floods training logs on automatic-registration buildsCó thể đã có người làm @ksivaman đã nhận hôm nay. Đang mở
Độ khó 4/5 3-5 ngày Mức phù hợp với người mới 50/100
NVIDIA/TransformerEngine#3645 · 1 bình luận · 1 người được giao ·
Maintainer thường phản hồi trong vòng 2 ngày
Tất cả issue của NVIDIA/TransformerEngine
Issue tương tự
-
Độ khó 2/5 1-3 giờ Mức phù hợp với người mới 70/100
NVIDIA/earth2studio#1241 ·
Maintainer thường phản hồi trong vòng 3 ngày
-
docs(types): update the collection binding note now that typed collections shipped in pycubrid 1.9.0Đang mởdocumentation priority: low size: S
Độ khó 2/5 1-3 giờ Mức phù hợp với người mới 75/100
cubrid-lab/sqlalchemy-cubrid#768 ·
Maintainer thường phản hồi trong vòng 1 ngày
-
--csv-bom was never wired up: PR #850 added an unused helper parameter, so #846 is not fixedĐang mởbug help wanted
Độ khó 2/5 1-3 giờ Mức phù hợp với người mới 75/100
Maintainer thường phản hồi trong vòng 1 ngày
-
Broken link in index.rstĐang mởdocumentation
Độ khó 1/5 Dưới một giờ Mức phù hợp với người mới 65/100
ansys/pydpf-core#3547 ·
Maintainer thường phản hồi trong vòng 1 ngày
-
good first issue
Độ khó 2/5 1-3 giờ Mức phù hợp với người mới 78/100
OktoLabsAI/okto-pulse#114 ·
Maintainer thường phản hồi trong vòng 1 ngày