Hacktoberfest 2026: những issue maintainer đã đánh dấu cho tháng Mười, đang mở và phù hợp người mới. Xem issue Hacktoberfest

FP8 DPA future token leakage

Đang mở
#3,248 1 bình luận 0 reaction 0 người được giao Xem trên GitHub

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
Công nghệ
python, pytorch

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ả

attention bug

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

  1. TE selects rowwise Q/K but columnwise V for forward MXFP8 attention, so V scales reduce along S: utils.py lines 2693–2727.
  2. TE passes V's columnwise payload and columnwise scale factors to the fused backend: fused_attn_fp8.cu lines 1105–1123.
  3. 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.h lines 326–413.
  4. cuDNN dequantizes V with block size {1, 32} immediately before BMM2/PV: scaled_dot_product_flash_attention.h lines 1026–1046.
  5. The causal mask is applied to the QK score tensor inside SDPA, after V has already been quantized externally: scaled_dot_product_flash_attention.h lines 860–884.
  6. 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.py lines 2777–2813.
  7. Current scaling computes max_fp8 / amax as a full FP32 value by default and clears the mantissa only when power-of-two scaling is explicitly enabled: recipe_common.cuh lines 14–48 and Float8CurrentScaling defaults 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

Bắt đầu từ đâu

  1. Đọc hết issue, rồi đọc hướng dẫn đóng góp của dự án.
  2. 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.
  3. Fork repository và làm thay đổi trên một nhánh.
  4. Mở pull request có tham chiếu số hiệu của issue.

Issue khác của NVIDIA/TransformerEngine

Tất cả issue của NVIDIA/TransformerEngine

Issue tương tự

Thêm issue về Python

Nhận issue mới trong hộp thư của bạn

Bản tóm tắt ngắn những issue GitHub phù hợp với người mới.