4678 字
23 分钟
DeepGEMM FP4 Indexer Kernel 实现解析

1. 背景:Lightning Indexer 与 FP4 Kernel#

DeepSeek V3.2 引入了 Lightning Indexer 机制来加速推理中的 token 检索。其核心操作可概括为:

给定 Q (query) 和全局 KV cache,计算每个 query token 与各 KV token 之间的加权 ReLU 注意力分数。

DeepGEMM 中负责这一操作的 kernel 即为 FP4 MQA Logits Kernel(对外 API 为 fp8_fp4_mqa_logits / fp8_fp4_paged_mqa_logits),在 README 中被称为 FP4 Indexer

Comment: 命名为 “indexer” 暗示这并非标准的 self-attention,而是用一种更简单的 scoring function 快速筛选候选 token,充当检索索引的角色。

本文仅关注 FP4 Indexer Kernel 的实现细节,不涉及其他 GEMM 变体(如 BF16 GEMM、Mega MoE 等)。


2. 两种模式:非分页与分页#

FP4 Indexer 提供两个 kernel 入口:

Kernel适用阶段源文件
sm100_fp4_mqa_logitsPrefill(KV cache 连续存放)sm100_fp4_mqa_logits.cuh
sm100_fp4_paged_mqa_logitsDecode(KV cache 分页存放)sm100_fp4_paged_mqa_logits.cuh

两者的计算核心完全相同;差异仅在于 KV cache 的寻址方式:非分页版使用 cu_seq_len_k_start/end 直驱索引,分页版使用 block_table + scheduler 间接寻址。


3. 输入张量的数据类型与形状#

3.1 Query (Q)#

属性
数据类型packed FP4cutlass::float_e2m1_t,每字节打包 2 个 FP4 元素)
形状 (非分页)[seq_len, num_heads, head_dim]
形状 (分页)[batch_size * next_n, num_heads, head_dim]
num_heads32 或 64
head_dim必须为 128(64B swizzle 对齐约束)

Comment: head_dim 写死为 128 是因为 UMMA 的 SWIZZLE_64B 布局需要 head_dim/2 的 8 倍对齐。这恰好与 DeepSeek V3.2 的 head dimension 一致。

3.2 Q Scale Factor (SF_Q)#

属性
数据类型UE8M0(4个 uint8 packed 进一个 int32
形状[seq_len, num_heads]
内存布局contiguous,TMA 搬运时对齐到 ceil_div(BLOCK_Q * num_heads, 128) 的元素数

对于 FP4 数据,每个 128 元素组共享一个 scale factor(block quantization),因此 SF_Q 的实际有效元素数为 BLOCK_Q * num_heads,但 TMA padding 会补到 128 的倍数。

3.3 KV Cache#

属性
数据类型packed FP4
形状 (非分页)[seq_len_kv, head_dim]
形状 (分页)[num_pages, page_size, head_dim]
head_dim128

KV cache 是 “headless” 的——所有 head 共享同一份 KV(MQA 特性)。

3.4 KV Scale Factor (SF_KV)#

属性
数据类型UE8M0(packed 进 int32
形状 (非分页)[seq_len_kv]
形状 (分页)[num_pages * page_size]
对齐要求16 字节 OOB safe,实际对齐到 ceil_div(BLOCK_KV, 128) 元素

3.5 Weights#

属性
数据类型float (FP32)
形状[seq_len, num_heads]
布局第 1 维 stride 必须为 1(contiguous on head dim)

Weights 在 Weighted ReLU 步骤中使用,每个 (token, head) 对应一个标量权重。

3.6 输出 Logits#

属性
数据类型用户指定(BF16 / FP32 / FP16 等)
形状 (标准)[seq_len, seq_len_kv]
形状 (compressed)[seq_len, max_seqlen_k]
对齐stride 对齐到 ceil_div(seq_len_kv + BLOCK_KV, 8) × 8

max_seqlen_k 非零时使用 compressed 格式,此时 logits 数组被 compact,且 kernel 不会清理未填充位置(clean_logits = false)。


4. 计算定义#

对于 query token i 和 KV token j,FP4 Indexer 的计算逻辑是:

# FP4 scale 解码与精度提升由 tensor core (FP19) 在内部流水线完成
# kv_j: [head_dim],tensor core 内以 FP19 精度参与乘法
# q_i: [num_heads, head_dim],tensor core 内以 FP19 精度参与乘法
# weights_i: [num_heads]
out_ij_h = q_i[h, :] @ kv_j # 内积 → 标量 (per head)
out_ij_h = relu(out_ij_h) * weights_i[h] # Weighted ReLU
out[i, j] = sum_h(out_ij_h) # reduce over heads → 标量

最终输出 out[i, j] 是跨所有 head 的加权 ReLU 内积之和。

Comment: 这本质上是一个等价于 sum(relu(QK^T) ∘ W) 的操作。与标准 attention 的区别在于:(1) 没有 softmax,(2) ReLU 替代了指数运算,(3) 权重矩阵 W 是可学习的 per-head 标量而非 QK 内积的产物。

4.1 中间数据类型追踪#

理解这条计算路径上的数据类型演变,是读懂 kernel 实现的前提。下面从输入端到输出端,逐步追踪每一步的数据类型与精度。

Step 0: 存储格式 → 逻辑值#

张量存储类型每元素字节数Packing逻辑值(经 SF 还原后)
Qfloat_e2m1_t (E2M1)0.5(2 元素/字节)packed FP4≥FP16(tensor core 内部 FP19)
KVfloat_e2m1_t (E2M1)0.5(2 元素/字节)packed FP4≥FP16(tensor core 内部 FP19)
SF_QUE8M00.25(4 元素/int32packed int322^sf,block size 128
SF_KVUE8M00.25(4 元素/int32packed int322^sf,block size 128
Weightsfloat4FP32

FP4 的元素格式是 E2M1(1-bit sign + 2-bit exponent + 1-bit mantissa),可表示的规格化范围为 ±[0.5, 1.75],再乘以 2^sf 的 block scale 恢复为实际值。需要特别说明的是,SM100 对 MXFP4 (Microscaling FP4) 的支持是一种混合精度方案

  • 存储位宽:4-bit(E2M1),这是带宽和显存占用的决定因素
  • Block scale 位宽:8-bit(UE8M0),每 128 元素共享一个 scale
  • Tensor core 内部算术精度:乘法和累加在更高精度下完成

这里的 “FP4 native support” 指的是硬件理解 E2M1 编码和 UE8M0 block scale 的语义,无需软件显式反量化——tensor core 从 SMEM 读到 FP4 字节流后,在内部流水线中自动完成 scale 解码与精度提升,然后执行乘法。中间的反量化值自始至终不存在于通用寄存器(GPREG)中,整个过程是 tensor core 指令(tcgen05.mma.kind::mxf4.block_scale)内部封闭的硬件行为。

至于 tensor core 内部的乘法精度,这里有一个容易忽略的细节:UE8M0 提供的是 8-bit unsigned exponent(范围 2^0 ~ 2^255),而 FP16 仅有 5-bit exponent(范围 2^-14 ~ 2^15)。如果内部乘法使用 FP16,则当 block scale 较大时,反量化后的值将立即溢出——这显然不可接受。因此 Blackwell tensor core 实际使用的内部格式是 FP19(1-bit sign + 8-bit exponent + 10-bit mantissa),其 8-bit exponent 恰好覆盖 UE8M0 的全范围而无需额外钳位。累加器则继续使用 FP32。

Comment: FP19 之于 MXFP4/MXFP6/MXFP8,类似于 FP32 之于 FP16/BF16/FP8——它是 tensor core 内部的”通用运算精度”。8-bit exponent 的设计使得它既能无损容纳 UE8M0 的全范围 scale,又比 FP32 节省了 13-bit mantissa(反正源头只有 4-bit,精度已定)。此外,“中间值不在寄存器上”这一点也对性能有直接影响:tensor core 的输入直接来自 SMEM 和 TMEM,无需将反量化结果回写 GPREG 再读入,省去了一个完整的数据往返。Q 的 128 元素 block size 恰好等于 UMMA_N = BLOCK_Q × num_heads,这意味着每个 UMMA tile 内部共享唯一的 SF_Q 和 SF_KV 组合,避免了 tile 内 SF 切换的开销。

Step 1: TMA 搬运 → SMEM#

数据搬运指令SMEM 布局SMEM 元素类型
QSM90_TMA_LOAD_2DSWIZZLE_64B, K-majoruint8(packed FP4 视为字节流)
KVSM90_TMA_LOAD_2DSWIZZLE_64B, K-majoruint8(packed FP4 视为字节流)
SF_Qtma::copy (1D)MN-major(TMA 后转置写入)int32(packed UE8M0)
SF_KVtma::copy (1D)MN-majorint32(packed UE8M0)
Weightstma::copy (1D)MN-majorfloat

SMEM 中的数据保持其原始存储类型。Q 和 KV 的 64B swizzle 布局是 UMMA 指令的对齐要求,数据本身不做类型转换。

Step 2: UMMA 发射 → TMEM 累加器#

UMMA 指令描述符中声明了 A/B 矩阵为 float_e2m1_t,累加器类型为 float。该指令在 SM100 tensor core 上的执行流程为:

  1. 从 SMEM 读取 packed FP4(SWIZZLE_64B 解交织后得到 E2M1 元素字节流)
  2. 从 TMEM 读取对应的 SF 值(UE8M0,已由 UTCCP 搬运到位)
  3. Tensor core 内部自动将 E2M1 解码并与 block scale 组合,提升到内部运算精度后执行乘法——整个过程在 tensor core 流水线内完成,无显式反量化步骤
  4. 矩阵乘累加:A(K-major) × B(K-major) → C,累加器为 float(FP32)
  5. 结果写入 TMEM 的累加器列
阶段数据类型
SMEM 读入 (A/B)float_e2m1_t(4-bit,packed 字节流)
SF 读入 (TMEM)UE8M0(硬件读取后解释为 2^sf
Tensor core 内部乘法E2M1 × 2^UE8M0 → FP19(1+8+10),乘法不在 4-bit 精度下执行;反量化后的值仅在 tensor core 数据通路中流动,不经过 GPREG
累加器 (TMEM)float (FP32)

Comment: MXFP4 的关键特性在于 SF 存储在 TMEM 而非 SMEM 中。与 FP8 不同(SF 可作为普通的 SMEM scale vector 传入),FP4 的 kind::mxf4.block_scale 指令要求 SF 必须驻留在 TMEM 内——这解释了为什么 kernel 需要先通过 UTCCP 将 SF_KV 搬到 TMEM 再做 UMMA。TMEM 作为 tensor core 的”近存”,SF 访问延迟极低,使得 block scale 能在每个 UMMA K-step 无开销地切换。

Step 3: TMEM → Math Registers#

Math warp 通过 SM100_TMEM_LOAD_32dp32b32x 指令将累加器从 TMEM 读入寄存器。对于 num_heads = 64,每次加载 32 个 float:

tmem_load(Int<32>{}, tmem_addr, accum);
tmem_load(Int<32>{}, tmem_addr + 32, accum + 32);
阶段数据类型
TMEM 累加器float (FP32)
寄存器 accum[]float (FP32)

加载后需要 fence_view_async_tmem_load() 保证 TMEM 读可见。

Step 4: ReLU + Weighted Sum(Math 阶段)#

这是唯一在常规寄存器中完成全部计算的步骤:

// Type of each operand:
sum_0, sum_1 → float2 (initialized with {0.0f, 0.0f})
accum[j] → float // ReLU input
fmaxf(...) → float // ReLU output
weights[i][j]→ float // from SMEM, via ld_shared
__ffma2_rn → float2 // FMA on two packed floats at once
__fadd2_rn → float2 // pairwise combine
sum.x + sum.y → float // final scalar reduction
操作输入类型输出类型精度
fmaxf(accum[j], 0)floatfloatFP32
weights[i][j] (from SMEM)floatfloatFP32
__ffma2_rnfloat2 × float2 + float2float2FP32, round-to-nearest-even
__fadd2_rnfloat2 + float2float2FP32, round-to-nearest-even
scalar add (sum.x + sum.y)floatfloatFP32

__ffma2_rn 是一条 PTX 级指令,在一个 cycle 内完成两组 a.x*b.x + c.xa.y*b.y + c.y,吞吐是两次标量 FMA 的两倍。四个 head 为一轮(两个 head 归入 sum_0,两个 head 归入 sum_1),最终 pair-wise 合并为单个标量。

Step 5: 写回 Global Memory#

auto result = static_cast<logits_dtype_t>(sum.x + sum.y);
logits[q_offset + kv_offset] = result;
阶段数据类型
归约结果(寄存器)float (FP32)
static_cast<logits_dtype_t>截断/舍入为 BF16、FP16 或保持 FP32
Global memory storeBF16 / FP16 / FP32(调用方可选)

logits_dtype_t 为 BF16 时,static_cast 将 FP32 的低 16 位丢弃(round-to-nearest-even 由硬件/编译器实现)。这是整个流程中唯一的精度损失点——从 FP32 累加器截断到目标输出类型。

数据类型流转总图#

Q (packed FP4, ½B/el) ──┐ Weights (FP32, 4B/el)
SF_Q (UE8M0, ¼B/el) ────┤ │
├─ UMMA: E2M1×2^UE8M0 ── FP19 mul (tensor core 内部, 不经过 GPREG) ── float acc ── TMEM
KV (packed FP4, ½B/el) ──┤ │
SF_KV (UE8M0, ¼B/el) ────┘ │
TMEM float ── load ── registers (float)
┌────────────────────┤
│ fmaxf (ReLU) │ ld_shared weights
▼ float ▼ float
__ffma2_rn ── float2 ── fadd2 ── float
static_cast ──▼
BF16 / FP16 / FP32 → Global Memory

Comment: 需要特别指出的是,MXFP4 的 block quantization 并非完全无损——128 个元素共享同一个 SF,意味着该 block 内的动态范围被限制为 ±1.75 × 2^sf。精度损失发生在量化阶段(FP32 → E2M1 + SF),而非 tensor core 内部运算。对于 MQA logits 这种 scoring 任务,量化损失通常不敏感于最终排序结果。Weights 保持 FP32 则是因为它在 ReLU 后直接参与乘法,若降低精度会引入额外的量化误差。


5. 核心实现:非分页版 sm100_fp4_mqa_logits#

5.1 整体架构:Warp 特殊化#

Kernel 采用 warp specialization,将线程划分为两类:

Warp 类型线程数职责
TMA Warp #01 warp (32 threads)TMA 异步发射 Q、SF_Q、Weights 的搬运请求
TMA Warp #11 warp (32 threads)TMA 异步发射 KV、SF_KV 的搬运请求
UMMA Warp1 warp (32 threads)SF 转置 (UTCCP) + 发射 UMMA FP4 指令
Math Warp GroupskNumMathWarpGroups × 4 warps从 TMEM 读取累加器、执行 Weighted ReLU + reduce、写回 global memory
Idle Warp (Specialized #3)1 warp寄存器重配后闲置

总线程数 = 128 (specialized) + kNumMathThreads (math threads),math threads 必须是 128 的倍数。对于典型配置(num_heads=64, head_dim=128),使用 128 + 256 = 384 线程,即 2 个 math warp group。

Comment: 这种 warp 特殊化设计将 producer (TMA) 和 consumer (UMMA, Math) 的职责完全分离,各自使用不同的寄存器预算(specialized warps 只配 56 个寄存器,math warps 配 224 个)。这是 SM100 上实现高吞吐的关键——它允许编译器在较少的寄存器压力下生成更高效的代码。

5.2 TMA 多级流水线#

Kernel 实现了两条并行的 TMA 预取流水线:

(a) Q 流水线(TMA Warp #0)

每个 Q block 包含 BLOCK_Q 个 token(BLOCK_Q = 128 / num_heads,使得每次搬运 BLOCK_Q * num_heads = 128 个元素对齐 UMMA_N)。TMA 一次搬运三份数据:

  • Q 数据:2D copy,[block_q * num_heads, head_dim] → SMEM
  • Q Scale Factor:2D copy,[num_heads, block_q] → SMEM (转置写入)
  • Weights:2D copy,[num_heads, block_q] → SMEM (转置写入)

Q 流水线级数 kNumQStages = 3,使用 transaction barrier 保证 TMA 写入完成后再由 consumer 读取。

(b) KV 流水线(TMA Warp #1)

每个 KV block 包含 BLOCK_KV 个 token(BLOCK_KV = kNumMathWarpGroups * UMMA_M,通常为 2 * 128 = 256)。TMA 搬运:

  • KV 数据:2D copy,[block_kv, head_dim] → SMEM
  • KV Scale Factor:2D copy,[block_kv] → SMEM

KV 流水线级数 kNumKVStages = 3

5.3 Scheduler:Range 边界计算#

对于非分页版,每个 Q block 通过 load_schedule() 计算其对应的 KV range:

for (uint32_t i = 0; i < BLOCK_Q; ++ i) {
auto row_idx = min(q_idx * BLOCK_Q + i, seq_len - 1);
seq_k_start[i] = min(cu_seq_len_k_start[row_idx], seq_len_kv);
seq_k_end[i] = min(cu_seq_len_k_end[row_idx], seq_len_kv);
start = min(start, seq_k_start[i]);
end = max(end, seq_k_end[i]);
}
start = start / 4 * 4; // TMA alignment for SF KV

它求出该 Q block 中所有 token 的 KV 访问区间并集的最小/最大值,形成批量 KV 搬运。随后 start 按 4 对齐以满足 SF TMA 的 16-byte 对齐要求。

5.4 UMMA 阶段:FP4 矩阵乘与 SF 处理#

UMMA Warp 负责从 SMEM 发射 tcgen05.mma.kind::mxf4.block_scale 指令。

5.4.1 Tile 尺寸#

维度
UMMA_M128
UMMA_NBLOCK_Q * num_heads (= 128)
UMMA_K64
K 步数head_dim / UMMA_K = 128 / 64 = 2

5.4.2 SF 搬运到 TMEM:UTCCP 转置#

Q Scale Factor 在 SMEM 中按 [num_heads, block_q] 布局存放(TMA 写入时 transpose)。在 UMMA 发射前,需要通过 UTCCP (Unified Tensor Core Copy) 将其搬运到 TMEM(Tensor Memory):

// Warp-level transpose via shared memory (UTCCP requires K-major layout)
utccp_required_smem_warp_transpose(smem_ptr);
cute::SM100_UTCCP_4x32dp128bit_1cta::copy(sf_desc, tmem_col);

UTCCP 要求数据在 SMEM 中为 K-major layout,而 TMA 写入的是 MN-major,因此需要一个 warp-level 的 SMEM 转置步骤(通过 ld_shared / st_shared + XOR lane swapping 实现)。

SF_Q 搬入 TMEM 后位于 tmem_col = kTmemStartColOfSFQ;SF_KV 搬入后位于 kTmemStartColOfSFKV

5.4.3 发射 UMMA#

auto runtime_instr_desc = make_runtime_instr_desc_with_sf_id(
instr_desc, k * 2, k * 2); // SF ids for A and B
auto a_desc = make_smem_desc(SWIZZLE_64B,
smem_kv + i * UMMA_M * (head_dim/2) + k * UMMA_K/2,
8 * (head_dim/2), 0);
auto b_desc = make_smem_desc(SWIZZLE_64B,
smem_q + k * UMMA_K/2,
8 * (head_dim/2), 0);
SM100_MMA_MXF4_SS::fma(a_desc, b_desc, tmem_addr, k,
runtime_instr_desc,
kTmemStartColOfSFKV + i * 4, kTmemStartColOfSFQ);

其中:

  • SM100_MMA_MXF4_SS::fma 封装了 PTX 指令 tcgen05.mma.cta_group::1.kind::mxf4.block_scale.block32(CUDA ≥ 12.9)或 kind::mxf4.block_scale.scale_vec::2X(旧版本)。
  • A 矩阵是 KV(K-major, SWIZZLE_64B),B 矩阵是 Q(K-major, SWIZZLE_64B)。
  • runtime_instr_desc 中的 a_sf_id / b_sf_id 分别指向 TMEM 中对应 K 步的 SF 列位置(A→SF_KV, B→SF_Q)。
  • 累加结果写入 TMEM 地址 tmem_addr = tmem_stage_idx * UMMA_N

UMMA 完成后通过 tcgen05.commit + mbarrier 通知 math warps。

5.4.4 TMEM 三级流水线#

与 Q/KV 类似,TMEM 也使用 3 级流水线(kNumTmemStages = 3)。UMMA warp 循环填充 TMEM stage,math warp groups 轮流消费。每个 stage 的容量为 UMMA_N = BLOCK_Q * num_heads = 128 列。

TMEM 总分配列数包括:

kNumTmemCols = align(
kNumAccumTmemCols // BLOCK_Q * num_heads * 3 stages
+ kNumSFQ / 32 // SF_Q columns
+ kNumSFKV / 32, // SF_KV columns
到最近的 2 的幂
)

Comment: 在 Blackwell 上,每个 SM 只有 512 列 TMEM(约 64KB),3 级流水线 + 两份 SF 存储,256 列左右,这在容量预算内。

5.5 Math 阶段:Weighted ReLU 归约#

Math warp groups 从 TMEM 读取累加器,执行后处理:

for (uint32_t i = 0; i < BLOCK_Q; ++ i) {
// Step 1: 从 TMEM 加载 accum[h] for h in [0, num_heads)
tmem_load(Int<num_heads/2>{}, tmem_addr, accum);
tmem_load(Int<num_heads/2>{}, tmem_addr + num_heads/2, accum + num_heads/2);
// Step 2: Weighted ReLU + pairwise reduction
const auto transform = [&](uint32_t j, float2 sum) {
auto a = make_float2(fmaxf(accum[j], 0), fmaxf(accum[j+1], 0)); // ReLU
auto b = make_float2(weights[i][j], weights[i][j+1]); // weights
return __ffma2_rn(a, b, sum); // FMA
};
// Pairwise: 每次处理 2 个 head,4 轮覆盖所有 head
for (uint32_t j = 0; j < num_heads; j += 4) {
sum_0 = transform(j, sum_0);
sum_1 = transform(j + 2, sum_1);
}
auto sum = __fadd2_rn(sum_0, sum_1);
auto result = static_cast<logits_dtype_t>(sum.x + sum.y);
// Step 3: 写回 global memory
logits[q_offset + kv_offset] = result;
}

几个值得注意的优化:

  • TMEM 加载粒度:每次加载 num_heads/2 个 float,对于 num_heads=64 就是 32 个 float,使用 SM100_TMEM_LOAD_32dp32b32x 指令。
  • __ffma2_rn:双路 FMA 指令,同时计算 a.x * b.x + c.xa.y * b.y + c.y,比两次标量 FMA 吞吐更高。
  • Pairwise reduction:将 num_heads 分成两个流(sum_0, sum_1),利用 float2 做向量化累加。
  • Warp-sync 调度:每处理完一个 BLOCK_Q 中的 token,执行 __syncwarp() 保证所有 lane 的 global store 完成,避免后续 barrier arrive 时出现 data race。
  • TMEM release 时机:当且仅当处理完 BLOCK_Q 的最后一个 token (i == BLOCK_Q - 1) 时 arrive empty barrier,最大化 TMEM 的利用窗口。

5.6 共享内存布局#

Kernel 的 shared memory 布局按如下顺序排列:

[ Q stage 0 | SMEM_Q_SIZE_PER_STAGE ]
[ Q stage 1 | SMEM_Q_SIZE_PER_STAGE ]
[ Q stage 2 | SMEM_Q_SIZE_PER_STAGE ]
[ KV stage 0 | SMEM_KV_SIZE_PER_STAGE ]
[ KV stage 1 | SMEM_KV_SIZE_PER_STAGE ]
[ KV stage 2 | SMEM_KV_SIZE_PER_STAGE ]
[ SF_Q stages 0-2 | 3 × SMEM_SF_Q_SIZE_PER_STAGE ]
[ SF_KV stages 0-2 | 3 × SMEM_SF_KV_SIZE_PER_STAGE ]
[ Weights stages 0-2 | 3 × SMEM_WEIGHT_SIZE_PER_STAGE ]
[ Transaction Barriers | (Q² + KV² + TMEM²) × 8 bytes ]
[ TMEM pointer | 4 bytes ]

所有 Q/KV buffer 对齐到 8 * (head_dim/2) = 512 bytes 以满足 TMA 64B swizzle 的要求。

5.7 Compressed Logits 模式#

kIsCompressedLogits = true(即 max_seqlen_k > 0),写回逻辑变为:

if (seq_k_start[i] <= kv_offset && kv_offset < seq_k_end[i])
logits[q_offset + kv_offset - seq_k_start[i]] = result;

每个 query token 的 logits 被 compact 到 [max_seqlen_k] 的空间内,且 kernel 不再负责清零无效位置——这由外部调用方控制。

Comment: 这种 compressed 格式在 indexer 中很有意义,因为每个 query token 实际需要的 KV 上下文窗口通常远小于 seq_len_kv,压缩存储可节省大量显存和带宽。


6. 分页版 sm100_fp4_paged_mqa_logits 的差异#

分页版的核心计算路径与非分页版一致。以下仅列出差异点。

6.1 新增参数:kNextN 与 Atom 粒度#

分页版引入 kNextN 参数,表示 batch 中每个序列同时处理的 query token 数。为了与 UMMA_N 对齐,将 kNextN 拆分为 atom:

  • kNextNAtom = (kNextN >= 2) ? 2 : 1
  • 每个 atom 最多处理 2 个 token 的 Q

这意味着 UMMA_N = kNextNAtom * num_heads,每次 UMMA 同时计算最多 2 个 query token 对同一 KV block 的内积。

6.2 Scheduler:smxx_paged_mqa_logits_metadata#

分页版使用一个独立的 metadata kernel 预先计算每个 SM 的调度起点:

  • 输入:context_lensindices(varlen 模式)
  • 输出:schedule_metadata[sm_idx * 2] = {q_atom_idx, kv_split_idx}

主 kernel 中通过 PagedMQALogitsScheduler 类读取 metadata,以 fetch_next_task() 逐任务推进。每次推进时:

  • 非 varlen 模式:atom 内 token 共享 context_len,只在 atom 边界刷新 num_kv_blocks
  • varlen 模式:每次推进后都重新计算 num_kv_blocks,并判断 is_paired_atom(两个连续 token 是否属于同一序列)

6.3 Block Table 寻址#

KV block 索引通过 block_table[atom_to_block_table_row(q_atom_idx)][kv_block_idx] 获取,其中 atom_to_block_table_row 将 atom 映射回原始 batch 行。

鉴于用户要求分页逻辑简略带过,此处不再展开。


7. 关键约束与设计权衡总结#

约束来源影响
head_dim = 128UMMA SWIZZLE_64B 对齐仅支持 128-dim head
num_heads ∈ {32, 64}BLOCK_Q * num_heads = 128 整除支持常用配置
SM100 only依赖 UMMA + TMEM + UTCCP仅 Blackwell GPU
CUDA ≥ 12.9 最佳mxf4.block_scale.block32 指令从 12.9 开始为直接模式旧版使用 scale_vec::2X fallback
Redundant writesMath warp 每个线程写全部 token(冗余写)源码标注 TODO 优化
Bank conflicts on weightsld_shared[i][j] 模式读取源码标注 TODO 优化

8. 调用示例(Python API)#

# 非分页版 (prefill)
logits = deep_gemm.fp8_fp4_mqa_logits(
q=(q_fp4, q_sf), # Q 数据 + SF
kv=(kv_fp4, kv_sf), # KV 数据 + SF
weights=weights, # [seq_len, num_heads] float
cu_seq_len_k_start=start,
cu_seq_len_k_end=end,
clean_logits=True,
max_seqlen_k=0,
logits_dtype=torch.bfloat16
)
# 分页版 (decode)
logits = deep_gemm.fp8_fp4_paged_mqa_logits(
q=(q_fp4, q_sf),
fused_kv_cache=(kv_cache, kv_cache_sf),
weights=weights,
context_lens=context_lens, # 2D: [batch_size, next_n]
block_table=block_table,
indices=indices, # optional, varlen
logits_dtype=torch.bfloat16
)

9. 小结#

FP4 Indexer Kernel 是 DeepGEMM 中专为 DeepSeek V3.2 的 Lightning Indexer 设计的高性能 CUDA kernel。它在 SM100 (Blackwell) 架构上充分利用了 UMMA FP4 指令、TMA 异步搬运、TMEM 暂存和 warp 特殊化技术,将 Q × KV^T → ReLU → Weighted Sum 这条计算路径高度 pipeline 化。数据精度方面,Q 和 KV 均为 packed FP4 + UE8M0 block scale factor,weights 保持 FP32,输出类型可由调用方指定,使得在保证检索质量的同时最大化显存带宽利用率。

Comment: 从源码风格来看,这是一个高度 “hand-tuned” 的 kernel——warp 分工、TMEM 列精准分配、barrier 时序、SF 转置策略均经过仔细编排。对于 CUDA kernel 优化的学习者而言,这份不到 500 行的实现是研究 SM100 UMMA 编程范式的极佳样本。

DeepGEMM FP4 Indexer Kernel 实现解析
https://infra.simphoni.uk/posts/fp4-indexer-deepgemm/
作者
Jingze Xing
发布于
2026-05-22
许可协议
CC BY-NC-SA 4.0