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_logits | Prefill(KV cache 连续存放) | sm100_fp4_mqa_logits.cuh |
sm100_fp4_paged_mqa_logits | Decode(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 FP4(cutlass::float_e2m1_t,每字节打包 2 个 FP4 元素) |
| 形状 (非分页) | [seq_len, num_heads, head_dim] |
| 形状 (分页) | [batch_size * next_n, num_heads, head_dim] |
num_heads | 32 或 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_dim | 128 |
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 ReLUout[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 还原后) |
|---|---|---|---|---|
| Q | float_e2m1_t (E2M1) | 0.5(2 元素/字节) | packed FP4 | ≥FP16(tensor core 内部 FP19) |
| KV | float_e2m1_t (E2M1) | 0.5(2 元素/字节) | packed FP4 | ≥FP16(tensor core 内部 FP19) |
| SF_Q | UE8M0 | 0.25(4 元素/int32) | packed int32 | 2^sf,block size 128 |
| SF_KV | UE8M0 | 0.25(4 元素/int32) | packed int32 | 2^sf,block size 128 |
| Weights | float | 4 | 无 | FP32 |
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 元素类型 |
|---|---|---|---|
| Q | SM90_TMA_LOAD_2D | SWIZZLE_64B, K-major | uint8(packed FP4 视为字节流) |
| KV | SM90_TMA_LOAD_2D | SWIZZLE_64B, K-major | uint8(packed FP4 视为字节流) |
| SF_Q | tma::copy (1D) | MN-major(TMA 后转置写入) | int32(packed UE8M0) |
| SF_KV | tma::copy (1D) | MN-major | int32(packed UE8M0) |
| Weights | tma::copy (1D) | MN-major | float |
SMEM 中的数据保持其原始存储类型。Q 和 KV 的 64B swizzle 布局是 UMMA 指令的对齐要求,数据本身不做类型转换。
Step 2: UMMA 发射 → TMEM 累加器
UMMA 指令描述符中声明了 A/B 矩阵为 float_e2m1_t,累加器类型为 float。该指令在 SM100 tensor core 上的执行流程为:
- 从 SMEM 读取 packed FP4(
SWIZZLE_64B解交织后得到 E2M1 元素字节流) - 从 TMEM 读取对应的 SF 值(UE8M0,已由 UTCCP 搬运到位)
- Tensor core 内部自动将 E2M1 解码并与 block scale 组合,提升到内部运算精度后执行乘法——整个过程在 tensor core 流水线内完成,无显式反量化步骤
- 矩阵乘累加:A(K-major) × B(K-major) → C,累加器为
float(FP32) - 结果写入 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 inputfmaxf(...) → float // ReLU outputweights[i][j]→ float // from SMEM, via ld_shared__ffma2_rn → float2 // FMA on two packed floats at once__fadd2_rn → float2 // pairwise combinesum.x + sum.y → float // final scalar reduction| 操作 | 输入类型 | 输出类型 | 精度 |
|---|---|---|---|
fmaxf(accum[j], 0) | float | float | FP32 |
weights[i][j] (from SMEM) | float | float | FP32 |
__ffma2_rn | float2 × float2 + float2 | float2 | FP32, round-to-nearest-even |
__fadd2_rn | float2 + float2 | float2 | FP32, round-to-nearest-even |
scalar add (sum.x + sum.y) | float | float | FP32 |
__ffma2_rn 是一条 PTX 级指令,在一个 cycle 内完成两组 a.x*b.x + c.x 和 a.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 store | BF16 / 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 ── TMEMKV (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 MemoryComment: 需要特别指出的是,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 #0 | 1 warp (32 threads) | TMA 异步发射 Q、SF_Q、Weights 的搬运请求 |
| TMA Warp #1 | 1 warp (32 threads) | TMA 异步发射 KV、SF_KV 的搬运请求 |
| UMMA Warp | 1 warp (32 threads) | SF 转置 (UTCCP) + 发射 UMMA FP4 指令 |
| Math Warp Groups | kNumMathWarpGroups × 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_M | 128 |
| UMMA_N | BLOCK_Q * num_heads (= 128) |
| UMMA_K | 64 |
| 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 Bauto 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.x和a.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_lens、indices(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 = 128 | UMMA 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 writes | Math warp 每个线程写全部 token(冗余写) | 源码标注 TODO 优化 |
| Bank conflicts on weights | ld_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 编程范式的极佳样本。