Gated Attention Block (SM107)
Gated Attention Block (SM107)
This is an experimental API and subject to change.
Overview
The gated attention block is the first model-level FE-OSS API: a set of FROST CuTe-DSL kernels behind one
Python class, one workspace and one execute() call. It implements the gated attention sub-layer used by
Qwen3.5-style models (the API is named by op geometry, the model is provenance only):
Every stage is a FROST kernel: the two projections drive the shipped FROST GEMM engine (pinned by name in the graph’s ranked plan list), stage (4) drives the shipped FROST SDPA through its standalone adapter, and stages (2)+(3) and (5) are this block’s own kernels. There are no cuBLAS or cuDNN-backend call-outs.
Target: NVIDIA Rubin (SM107, compute capability 10.7) only. Every other architecture is declined with a
typed error (NotImplementedError), never served slowly.
Three precisions share the signature, and two fp4 modes ride the MXFP8 one (both are fields on MxQuantSpec,
so they are unspellable on the bf16 and per-tensor FP8 pipelines rather than declined):
Fusion knobs
Two optional fusions are constructor flags; each is a different compiled specialization behind the same
execute() signature, and both default to off.
fuse_norm_rope=Truefolds stages (2)+(3) into stage (1)‘s epilogue: Q/K tiles are normed and rotated on the fp32 accumulator and written once (needsinplace_qkv; inference only, no pre-norm Q/K is kept). Under FP8 / MXFP8 the same epilogue also quantizes: compact e4m3q8/k8/v8(and the MXFP8 scale factors) come straight out of the GEMM, so the quantize passes disappear.fuse_gate=Truefolds stage (5) into the SDPA epilogue:O *= sigmoid(GATE)after the dead-row select, with the gate tile TMA-staged by the load warp (inference only: no pre-gateOfor the backward).
With both on, the block is three launches: proj(+norm+RoPE[+quant]) -> sdpa(+gate) -> out_proj.
geometry.qk_norm=False runs RoPE-only Q/K on every path (norm weights are passed as None, no rstd is
produced, the backward drops the norm-weight gradients).
API
Geometry
GatedAttentionBlockGeometry (frozen dataclass; validate() raises ValueError):
Weights and tables
W_qkvg [N, d_model]withN = (2*H_q + 2*H_kv) * D, column blocksQ | GATE | K | V. Build it from separate projection weights withbuild_fused_qkvg_weight(w_q_gate, w_k, w_v, geometry, q_gate_layout="flat"); the block’s tile alignment isQKVG_TILE_ALIGN = 64columns.W_o [d_model, H_q * D].cos,sin[B, S, rope_dim]rotary tables in the activation dtype (rotate-half convention on the firstrope_dimdims of every head).w_q_norm,w_k_norm[D](bothNoneiffgeometry.qk_normisFalse).- MXFP8 only:
h_sfandw_qkvg_sf, the E8M0 scale factors ofhandW_qkvgin cuDNN’s F8_128x4 order (uint8orfloat8_e8m0fnu; byte counts fromcudnn.gated_attention_block.kernels.proj_gemm.sf_blob_bytes). - fp4 weights are packed e2m1, dtype
torch.float4_e2m1fn_x2, stored[N, K // 2]— two codes per byte along the contraction axis, LOW nibble = evenk. The block checks the STORAGE shape ([N, d_model // 2]forW_qkvg,[d_model, H_q * D // 2]forW_o); a logical[N, K]fp4 tensor, oruint8storage, is a typedValueError(torch can.view(torch.float4_e2m1fn_x2)packed bytes but cannot cast to fp4).- MXFP4
W_qkvg(MxQuantSpec.w_qkvg_dtype=torch.float4_e2m1fn_x2):w_qkvg_sfis UNCHANGED — the same E8M0 / 32 F8_128x4 blob overn_qkvg x d_modelas for e4m3 codes. - fp4
W_o(MxQuantSpec.o_fp4): its scale blobw_o_sfis in the SAME format asO—float8_e4m3fnscales per 16 forFp4Format.NVFP4,float8_e8m0fnuper 32 forFp4Format.MXFP4(either asuint8or as that dtype),sf_blob_bytes(d_model, H_q * D, block)bytes in F8_128x4 order (padded to whole 128-row x 4-block atoms; pad bytes, if any,0x00).
- MXFP4
Forward
Every appended argument (sample_w_o_sf, w_o_sf) sits at the end with a None default, so positional callers of the
bf16, FP8 and MXFP8 pipelines are unchanged; MxQuantSpec.o_fp4 and sample_w_o_sf must be given together (a typed
ValueError names the missing half).
execute() allocates nothing, reads nothing back to the host and converts nothing: every intermediate is a
strided view of the caller’s workspace, sized honestly by get_workspace_size(), so the call is CUDA-graph
friendly. All stages run on one launch stream (torch’s current stream, or current_stream), on both the FROST
GEMM route and the DSL stages. Dtype, layout and shape mismatches, an unsupported architecture, or a feature the
selected specialization cannot serve raise typed errors at construction (check_support) rather than at launch.
Quantization specs:
-
QuantSpec(descale_h, descale_w_qkvg, descale_w_o, scale_q, scale_k, scale_v, scale_o, dtype=torch.float8_e4m3fn)— static per-tensor scales for the FP8 pipeline. -
MxQuantSpec(descale_w_o, scale_o=1.0, dtype=torch.float8_e4m3fn, block_size=32, w_qkvg_dtype=torch.float8_e4m3fn, o_fp4=None)— the MXFP8 pipeline; the MXFP8 SDPA writes e4m3Ounscaled, so the fully fused path needsscale_o == 1.0. The two appended fields select the fp4 modes:w_qkvg_dtype:torch.float8_e4m3fn(default, MXFP8 x MXFP8) ortorch.float4_e2m1fn_x2(MXFP4W_qkvg, the mixed row); any other dtype is a typedNotImplementedError.o_fp4:None(default, per-tensor e4m3O) or anFp4Formatmember (anything else is aTypeError). Undero_fp4neither side of the out projection has a per-tensor scale — the quantizer writes block scales only andW_odequantizes throughw_o_sfin the MMA — soscale_oanddescale_w_omust both be1.0(typedValueErrorotherwise; a non-unit value would be silently dropped by the block-scale GEMM).
-
Fp4Format(enum; one member = codes x scale dtype x block, so an illegal pairing cannot be spelled):Properties:
block_size,sf_torch_dtype,sf_cudnn_dtype(andfmt_name, the kernel’s format key). Neither format carries a global (per-tensor) scale: an NVFP4 block’s e4m3 scale ismax(amax * fp32(1/6), 2^-9)(the e4m3 min-subnormal floor keeps an all-zero block’s scale NONZERO, so the encodex / scalestays finite), an MXFP4 block’s E8M0 scale is the power of two at or aboveamax * fp32(1/6)—fp32(1/6), not an exact/ 6: the kernel and the reference multiply by the fp32 constant, and the two differ by one ulp; a dead (fully masked or zero-length) row quantizes to codes0exactly. TheOcodes are round-to-nearest-even on the e2m1 grid, saturating at 6.
Backward
GatedAttentionBlockBwd(sample_dy, sample_saved, sample_w_qkvg, sample_w_q_norm, sample_w_k_norm, sample_cos, sample_sin, sample_w_o, geometry, *, recompute=RecomputePolicy.RECOMPUTE_QK_PRE, need_dh=True, need_dw_qkvg=True, need_dw_o=True, need_dw_norms=None) consumes the forward’s SavedForBackward(h, gate, o, lse, rstd_q, rstd_k, q_pre=None, k_pre=None). RecomputePolicy chooses between re-running stage (1) for the pre-norm
Q/K (RECOMPUTE_QK_PRE, the default) and reading them from the save set (SAVE_ALL); which input gradients are
wanted is fixed at build time because it decides which GEMMs exist. need_dw_norms=None follows
geometry.qk_norm; asking for norm-weight gradients under qk_norm=False is a typed decline. The backward is
bf16 / fp16 only.
Requirements and limits
- Rubin (SM107) only; cuDNN 9.x,
nvidia-cutlass-dsl >= 4.8.0.dev0(the Rubin arch names), torch. d_head = 256(the Rubin d256 SDPA flavor with the fused gate);d_model % 128 == 0under MXFP8.- FP8 / MXFP8 are inference only; the backward is bf16 / fp16.
- FP8: a dense (no-mask) sequence length must be a multiple of 128 unless the causal mask or a padding mask
covers the KV tail (the Rubin per-tensor FP8 SDPA contract); MXFP8: e4m3 codes only (e5m2 is a typed decline);
the fully fused MXFP8 path needs
scale_o == 1.0and, atB > 1,S % 128 == 0(a scale-factor atom is per sequence). fuse_gateandfuse_norm_ropeare inference-only specializations (no pre-gateO, no pre-norm Q/K).- fp4 (
MxQuantSpec.w_qkvg_dtype/o_fp4): MXFP8 pipeline only (unrepresentable onQuantSpec/ bf16); inference only (save_for_backwardis a typed decline, as for every quantized pipeline); no global per-tensor scale in either fp4 format (scale_o == descale_w_o == 1.0undero_fp4);d_head % (4 * block) == 0undero_fp4(whole 4-block scale words per head: 64 for NVFP4, 128 for MXFP4;d_head = 256passes both); the MXFP4W_qkvgruns on the unfused pipeline only —fuse_norm_ropewith an e2m1W_qkvgis a typedNotImplementedError(the fused MXFP8 projection fork is rendered for an e4m3 B).hstays e4m3 (an fp4his not served), and an fp4W_owith an e4m3Ois not a served pairing.
Performance
Whole block, B=1, h_q=32 h_kv=2 d=256 d_model=5120 (the 397B geometry), Rubin perf node (212 SMs, SM clock
locked at 2376 MHz), speedup over the same bf16 torch chain (median of 5 launch-interleaved rounds x 30 launches;
the bf16 FROST control pair stayed within 0.6 %). The fp4 modes (MXFP4 weights, NVFP4 / MXFP4 O) are not in these
tables: their perf-node measurement is pending, and no number is quoted until it exists.
Causal:
Dense (no mask):
Related
- The SDPA epilogue gate is also reachable through the graph API: an
sdpanode followed bysigmoidandmulpointwise nodes onOis served fused by the Rubin d256 FROST SDPA engines — see Attention, “Fused epilogue gate”. - How the block is composed (workspace, streams, fusion knobs, typed declines) and how to build the next one: Composing multi-kernel blocks in Python.
- The MLA sibling of the fused projection epilogue: GEMM + RoPE + MXFP8 Projection.
- Tests:
test/python/fe_api/gated_attention_block/(layout contract, reference oracle, end to end, FP8, MXFP8, fp4 weights / fp4 O (test_block_fp4.py,test_proj_gemm_fp4.py,test_quantize_fp4.py), per-stage kernels, stream ordering).