NVFP4 Attention QAT Backward
This is an experimental API and subject to change.
Overview
nvfp4_attention_qat_backward computes explicit Q, K, and V gradients for
scaled dot-product attention trained with NVFP4 fake quantization. It is a
Triton port of FastVideo’s attention QAT backward at commit
e9bbaca07d511b2ee7e16474dae6f923426223dc:
The operation fake-quantizes Q, K, and V to the NVFP4 E2M1 data format with an E4M3 scale for every 16 values, then immediately dequantizes them for the attention computation. The probability matrix follows two paths:
- dQ and dK use the unquantized softmax probability, implementing the straight-through estimator (STE).
- dV uses the NVFP4 fake-quantized probability.
The implementation launches four kernels: fused Q fake-quantization/delta preprocessing, K/V fake-quantization, dQ, and dK/dV. Causal backward skips fully masked tiles while retaining the elementwise mask on the diagonal. The local-scale conversion uses precise division so exact E2M1 midpoints preserve round-to-nearest-even; no tolerance relaxation is required. The production non-causal SM100 configuration uses 64 by 64 tiles; other supported Blackwell configurations use 32 by 32 tiles.
Installation
From a source checkout, install the CuTe DSL base dependencies, Triton, and the torch dependency group:
For a published wheel, install a CUDA-enabled torch build separately:
Triton 3.7 or newer is supported on Linux with Python 3.10 or newer for this API.
High-level wrapper
The result keys are dq_tensor, dk_tensor, and dv_tensor. Optional
preallocated tensors with those names can be passed to the wrapper. Pass a
cuda.CUstream as current_stream to order wrapper allocations and all
kernel launches on an explicit stream.
high_precision_o is not the probability-quantized user-visible QAT output.
It must be the matching softmax(Q_fake K_fake^T) @ V_fake value saved before
probability fake quantization. lse is the corresponding natural-log
log-sum-exp statistic. Supplying forward auxiliaries from a different
quantization recipe produces incorrect gradients.
Class API
Nvfp4AttentionQatBackward exposes explicit validation, compilation, and
execution. execute performs no allocations; the caller supplies contiguous
gradient buffers and a one-dimensional CUDA torch.uint8 workspace.
compile() materializes every shape- and architecture-specialized Triton
kernel without launching it. execute() then reuses those cached artifacts.
Tensor contract
All tensors must use contiguous, 16-byte-aligned BHSD storage and reside on one
CUDA device. softmax_scale defaults to 1 / sqrt(128) and must match the
forward pass.
Current support and limitations
- GPU: SM100, SM103, SM120, and SM121 Blackwell.
- Attention: MHA with equal query and KV head counts; head dimension 128.
- Sequence lengths: self-attention and non-causal cross-attention, including non-aligned tails.
- Causal mode: self-attention only.
- Dtype: BF16 activations and FP32 LSE.
- Not implemented: GQA/MQA, dropout, padding or packed variable-length sequences, bias, local masks, and deterministic-mode selection.