GEMM + RoPE + MXFP8 Projection (SM100)
GEMM + RoPE + MXFP8 Projection (SM100)
This is an experimental API and subject to change.
JAX support
Supports JAX arrays on both input paths (BF16 and MXFP8) with w_out_in=True (the [in, out] weight layout reaches the kernel through a transposed strided view, which has no row-major JAX equivalent and raises a clear error). The E8M0 scale inputs stay uint8 as with torch. Outputs are allocated as C-contiguous jnp arrays. The wrapper is eager only, on the CUDA legacy default stream: block_until_ready inputs, synchronize before reading outputs.
For jitted JAX programs use the jax.jit-compatible XLA custom-call entry point gemm_proj_rope_mxfp8_jax_sm100(x, w, cos, sin, x_scale=None, w_scale=None) (built on cudnn.jax.call; see gemm_amax.md “Using JAX arrays”): same contract as the wrapper with w_out_in=True, dispatching on x.dtype (bfloat16 → BF16 GEMM; float8_e4m3fn plus E8M0 scales → MXFP8 GEMM), returning (out_fp8_row, out_scales_row, out_fp8_col, out_scales_col) as fresh XLA-managed arrays — no manual synchronization needed, composes with jax.jit and CUDA graphs.
The API is compiled with --enable-tvm-ffi: raw framework tensors go straight to the compiled kernel (no per-call from_dlpack conversion), cutting per-launch CPU overhead roughly in half for torch callers as well.
Overview
Fused projection GEMM + per-head YARN RoPE + dual-direction MXFP8 quantize: a persistent dense GEMM on NVIDIA Blackwell GPUs (SM100+) that projects activations, applies the Megatron MLA-YARN rotary embedding to each attention head’s trailing rotary features, and MXFP8 (E4M3, block=32) quantizes the result in both the rowwise (D-direction) and columnwise (S-direction) layouts. Implemented with CUTLASS/CUTE.
Two sibling kernels implement the same operation, differing only in the GEMM input precision, and are selected by the dtype of x/w:
- BF16 input (
GemmProjRopeMxfp8Bf16InSm100):xandwarebfloat16; the GEMM runs in bf16. - MXFP8 input (
GemmProjRopeMxfp8Mxfp8InSm100):xandware pre-quantized MXFP8 (E4M3 codes + E8M0 rowwise block scales); the GEMM runs in MXFP8.
The output is dual-direction MXFP8 in both cases. Emitting both scale directions makes the output directly consumable by block-scaled matmuls that need either operand orientation. For example, in DeepSeek-V3 the rowwise output feeds the forward QK^T and the columnwise output feeds the backward dK = dS^T · Q on the cuDNN is_input_fp8 attention path.
- Inputs: activations
x, projection weightw, and bf16 rotary tablescos/sin; for the MXFP8-input path, the E8M0 block scalesx_scale/w_scaleas well. - Outputs: rowwise and columnwise MXFP8 data (
out_fp8_row,out_fp8_col) and their E8M0 scale factors (out_scales_row,out_scales_col).
The kernel is tuned for the DeepSeek-V3 Q up-projection shapes: NUM_HEADS=128, HEAD_DIM=192 (QK_NOPE=128 + QK_ROPE=64), MXFP8 BLOCK=32, tile TILE_M=128 (one head per CTA).
Shapes
-
Inputs
x:(tokens, Q_LORA)—tokens % TILE_M == 0. Dtypebfloat16(bf16 path) orfloat8_e4m3fn(MXFP8 path).w:(Q_LORA, NUM_HEADS·HEAD_DIM)whenw_out_in=False, or the transformer-engine-native transposed(NUM_HEADS·HEAD_DIM, Q_LORA)whenw_out_in=True. Same dtype asx.x_scale,w_scale(MXFP8 path only): E8M0 rowwise block scales,uint8.x_scaleis(tokens, Q_LORA // BLOCK).w_scalefollowsw’s layout (the wrapper transposes it alongsidew):(NUM_HEADS·HEAD_DIM, Q_LORA // BLOCK)whenw_out_in=True, or(Q_LORA // BLOCK, NUM_HEADS·HEAD_DIM)whenw_out_in=False.cos,sin:(tokens, QK_ROPE),bfloat16.
-
Outputs
out_fp8_row,out_fp8_col:(tokens, NUM_HEADS, HEAD_DIM)out_scales_row:(tokens, NUM_HEADS, HEAD_DIM // BLOCK)out_scales_col:(tokens // BLOCK, NUM_HEADS, HEAD_DIM)
Equations
Project and reshape per head, then apply the interleaved-in / halves-out YARN RoPE to the trailing QK_ROPE features of each head:
MXFP8 quantize with block size BLOCK=32, independently for each direction (E8M0 per-block scale, E4M3 data):
Diagram
API Usage
High-level wrapper (dtype-dispatch)
Selects the bf16- or mxfp8-input kernel by x/w dtype (which must match); pass x_scale/w_scale only for the MXFP8 path.
Class API — BF16 input
The class constructor defaults to w_out_in=False (w stored [in, out]); pass w_out_in=True for TE-native [out, in] weights. (The high-level wrapper defaults the other way, to w_out_in=True.)
Class API — MXFP8 input
Unlike the bf16 class, this class has no w_out_in parameter: w_code/w_scale must already be TE-native [out, in] = [N, K]. Use the high-level wrapper if your weight is [in, out] — it transposes the code and scale before constructing this API.
Parameters
Input/Output tensors
- Input x:
(tokens, Q_LORA); Dtypebfloat16(bf16 path) orfloat8_e4m3fn(MXFP8 path). - Input w:
(Q_LORA, NUM_HEADS·HEAD_DIM)(w_out_in=False) or(NUM_HEADS·HEAD_DIM, Q_LORA)(w_out_in=True); same dtype as x. - Input x_scale, w_scale (MXFP8 path only): Dtype
uint8(E8M0 rowwise block scales).x_scaleis(tokens, Q_LORA // BLOCK).w_scale’s shape depends onw_out_in(transposed withw):(NUM_HEADS·HEAD_DIM, Q_LORA // BLOCK)forw_out_in=True,(Q_LORA // BLOCK, NUM_HEADS·HEAD_DIM)forw_out_in=False. - Input cos, sin:
(tokens, QK_ROPE); Dtypebfloat16. - Output out_fp8_row / out_fp8_col:
(tokens, NUM_HEADS, HEAD_DIM); Dtypefloat8_e4m3fn. - Output out_scales_row:
(tokens, NUM_HEADS, HEAD_DIM // BLOCK); Dtypeuint8(E8M0). - Output out_scales_col:
(tokens // BLOCK, NUM_HEADS, HEAD_DIM); Dtypeuint8(E8M0).
Common parameters
w_out_in: bool— whetherwis stored[out, in](True) or[in, out](False). Wrapper default:True. On the bf16 path the kernel consumes both via the cutlass major mode (no transposed copy); on the MXFP8 path the wrapper transposes the code + scale for[in, out].x_scale,w_scale: Optional[Tensor]— required forfloat8_e4m3fninputs; must beNoneforbfloat16.x.dtype == w.dtypeis asserted.- CUDA stream (
current_streamin class API,streamin wrapper). Defaults to the current torch stream (required for CUDA-graph capture).
Wrapper return values
Returns a TupleDict with keys out_fp8_row, out_scales_row, out_fp8_col, out_scales_col. Tuple unpacking order is (out_fp8_row, out_scales_row, out_fp8_col, out_scales_col).
Support surface and constraints
Dtypes
- BF16 path:
x,w,cos,sinarebfloat16. - MXFP8 path:
x,warefloat8_e4m3fncodes;x_scale,w_scaleareuint8(E8M0);cos,sinarebfloat16. - Outputs (both paths):
out_fp8_row,out_fp8_colarefloat8_e4m3fn;out_scales_row,out_scales_colareuint8(E8M0). - No fp4 (e2m1) input on either path: the fused projection epilogues — this kernel and the
gated attention block’s
fuse_norm_ropefork twin alike — are rendered for ane4m3B operand. The gated block serves an MXFP4W_qkvgon its unfused pipeline only (the FROST GEMM’s mixed MXFP8 x MXFP4 block-scale row); asking for it withfuse_norm_rope=Trueis a typedNotImplementedError.
Shapes and divisibility
tokens % TILE_M == 0(TILE_M = 128); no tail handling.- The projected weight dimension must equal
NUM_HEADS·HEAD_DIM. On the MXFP8 path,Q_LORA % BLOCK == 0.
Environment
- Requires CUDA with SM100+ compute capability.
Source provenance
Integrated from the DeepSeek-V3 MLA fused Q up-projection kernel developed for Megatron-LM MXFP8 training (Blackwell / customte CUTLASS 4.4.1); originally added in the GEMM+RoPE+MXFP8 fusion commit. The BF16-GEMM and MXFP8-GEMM variants were consolidated into two input-precision kernel modules selected by input dtype:
python/cudnn/gemm/cutedsl/dense/proj_rope_mxfp8/gemm_proj_rope_mxfp8_bf16in.py— BF16-input kernel; also hosts the pure-PyTorch oraclegemm_proj_rope_mxfp8_reference(...).python/cudnn/gemm/cutedsl/dense/proj_rope_mxfp8/gemm_proj_rope_mxfp8_mxfp8in.py— MXFP8-input kernel.
The compiled-kernel lifecycle (check_support/compile/execute) lives in the APIBase classes in api.py; the wrapper caches the compiled objects (matching the sibling GEMM-fusion packages).
Installation
Requires the CuTeDSL dependencies, which ship with the package:
Usage examples
For usage examples, see test cases in test/python/fe_api/gemm/test_gemm_proj_rope_mxfp8.py.