submission 738050
pingp2574 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 572 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-738050?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:9ce1b6024c19a85928676d5c448e20515f10f616963462cc09da8249135f6b7b
license declaredunknown
license concludedunknown
authorspingp2574
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
4. _dq_bf16_buf: 反量化缓冲区(MXFP4 → bf16 时使用)persistent-kernel
获取或创建 MLA 解码内核所需的持久模式元数据(persistent-mode metadata)。Kernel source
submission.py572 lines
"""
Submission 实现:MLA (Multi-head Latent Attention) 解码阶段的注意力计算。
=== 背景知识 ===
这是 DeepSeek R1 模型中使用的 MLA 注意力机制的解码(decode)阶段实现。
MLA 的核心思想是:将传统的 Multi-Head Attention 中的 KV 缓存进行压缩,
使用一个低秩的隐空间表示(latent representation),从而大幅减少 KV 缓存的显存占用。
=== MLA 的 forward_absorb 路径 ===
在 DeepSeek R1 中,MLA 采用了 "forward_absorb"(前向吸收)策略:
- Q(Query)被投影到一个 576 维的吸收空间(absorbed space)
- KV 缓存也被压缩到同样的 576 维空间
- 576 = 512 (kv_lora_rank, 用于 V 输出) + 64 (qk_rope_head_dim, 用于旋转位置编码)
=== 数据流概览 ===
输入: q (total_q, 16, 576) bf16 + KV缓存 (total_kv, 1, 576) fp8
↓ ↓
量化为 FP8 已预量化为 FP8
↓ ↓
└──────── mla_decode_fwd ────────┘
↓
输出: attention output (total_q, 16, 512) bf16
=== 本文件的策略 ===
本实现使用 FP8 精度进行注意力计算(a8w8 模式:Q 和 KV 都用 FP8),
相比 bf16 在 AMD MI355X 上可获得约 2-3 倍的加速,精度损失可忽略。
同时使用动态 KV 分片策略(num_kv_splits),根据 batch_size * kv_seq_len
的总工作量自适应调整分片数量,在保证精度的同时优化性能。
"""
import torch
from task import input_t, output_t
# aiter 是 AMD 的 AI 推理加速库,提供了高性能的 MLA 注意力内核
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.utility.fp4_utils import (
mxfp4_to_f32,
e8m0_to_f32,
)
# ========================================================================
# DeepSeek R1 MLA 模型超参数(forward_absorb 路径)
# 这些值来自 DeepSeek-R1 的 config.json 配置文件
# ========================================================================
# 注意力头数:16 个 Query 头
NUM_HEADS = 16
# KV 头数:只有 1 个(MLA 本质上是 Multi-Query Attention,所有 Q 头共享一个 KV 头)
# 这是 MLA 节省显存的关键之一
NUM_KV_HEADS = 1
# KV LoRA 秩:KV 压缩后的隐空间维度为 512
# 原始的 KV 维度远大于 512,通过低秩投影压缩到 512 维
KV_LORA_RANK = 512
# QK 旋转位置编码(RoPE)维度:64
# RoPE 是一种相对位置编码,只作用于 QK 计算的一部分维度
QK_ROPE_HEAD_DIM = 64
# QK 头总维度 = KV LoRA 秩 + RoPE 维度 = 512 + 64 = 576
# 576 维中:前 512 维用于 V 输出,全部 576 维用于 QK 匹配
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
# V 头维度 = KV LoRA 秩 = 512
# 注意:V 只使用前 512 维,不包含 RoPE 的 64 维
V_HEAD_DIM = KV_LORA_RANK
# Softmax 缩放因子:1 / sqrt(d_k) = 1 / sqrt(576)
# 防止点积结果过大导致 softmax 梯度消失
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
# 页大小:PagedAttention 中的页大小为 1(即每个 KV token 独立存储)
PAGE_SIZE = 1
# FP8 数据类型(平台相关,由 aiter 库决定具体是 FP8 E4M3 还是 E5M2)
FP8_DTYPE = aiter_dtypes.fp8
# FP8 可表示的最大值,用于量化时的缩放计算
FP8_MAX = torch.finfo(FP8_DTYPE).max
class _BufferCache:
"""
GPU 显存缓冲区缓存(Buffer Cache)。
=== 为什么需要这个类? ===
在 LLM 推理的解码阶段,每次生成一个新 token 都需要调用注意力计算。
如果每次调用都重新分配 GPU 显存,会带来巨大的开销(CUDA malloc 很慢)。
这个类通过缓存和复用 GPU 张量来避免重复的显存分配:
- 当请求的张量形状与之前相同时,直接复用已分配的显存
- 当形状不同时,才重新分配
=== 缓存的缓冲区类型 ===
1. _kv_indices: KV 索引数组,用于 PagedAttention
2. _output_buf: 输出缓冲区,存储注意力计算结果
3. _meta: 持久模式的元数据(work_metadata 等),用于优化内核调度
4. _dq_bf16_buf: 反量化缓冲区(MXFP4 → bf16 时使用)
"""
# __slots__ 禁止动态添加属性,节省内存并提高访问速度
__slots__ = (
"_meta_key", "_meta",
"_kv_indices", "_kv_indices_len",
"_output_buf", "_output_key",
"_dq_bf16_buf", "_dq_bf16_key",
)
def __init__(self):
self._meta_key = None
self._meta = None
self._kv_indices = None
self._kv_indices_len = 0
self._output_buf = None
self._output_key = None
self._dq_bf16_buf = None
self._dq_bf16_key = None
def get_kv_indices(self, total_kv_len: int) -> torch.Tensor:
"""
获取 KV 缓存的索引数组 [0, 1, 2, ..., total_kv_len-1]。
在 PagedAttention 中,kv_indices 用于指定要访问的 KV 缓存页。
当 PAGE_SIZE=1 时,索引就是简单的连续整数序列。
缓存策略:如果已有的数组长度 >= 需要的长度,直接切片复用;
否则重新创建一个更大的数组。
"""
if self._kv_indices is None or self._kv_indices_len < total_kv_len:
self._kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
self._kv_indices_len = total_kv_len
return self._kv_indices[:total_kv_len]
def get_output(self, total_q: int, nq: int, dv: int) -> torch.Tensor:
"""
获取输出缓冲区,形状为 (total_q, nq, dv)。
参数:
total_q: 所有 batch 中 Q token 的总数
nq: 注意力头数 (16)
dv: V 头维度 (512)
返回: 未初始化的 bf16 张量(后续会被 mla_decode_fwd 写入结果)
"""
key = (total_q, nq, dv)
if self._output_key != key or self._output_buf is None:
self._output_buf = torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda")
self._output_key = key
return self._output_buf
def get_metadata(
self,
batch_size, max_q_len, nhead, nhead_kv,
q_dtype, kv_dtype,
qo_indptr, kv_indptr, kv_last_page_len,
num_kv_splits,
):
"""
获取或创建 MLA 解码内核所需的持久模式元数据(persistent-mode metadata)。
=== 什么是持久模式(Persistent Mode)? ===
在传统的 GPU 内核调用中,每次启动内核都需要重新计算调度信息。
而在解码阶段,连续的 token 生成之间,很多调度参数是不变的。
持久模式通过预先计算并缓存这些调度信息,避免重复计算,提高效率。
=== 元数据包含什么? ===
- work_metadata: 工作分配元数据(哪个 GPU 线程块处理哪个任务)
- workIndptr: 工作项的索引指针
- workInfoSet: 工作项的详细信息
- reduceIndptr: 归约操作的索引指针(用于 KV 分片后的结果合并)
- reduceFinalMap: 最终归约的映射表
- reducePartialMap: 部分归约的映射表
=== 缓存策略 ===
使用 (batch_size, max_q_len, nhead, nhead_kv, dtype, num_kv_splits) 作为缓存键。
如果键匹配,只重新填充元数据(不重新分配显存);
如果键不匹配,需要重新分配元数据缓冲区。
"""
key = (batch_size, max_q_len, nhead, nhead_kv, str(q_dtype), str(kv_dtype), num_kv_splits)
if self._meta_key == key and self._meta is not None:
(work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = self._meta
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
nhead // nhead_kv,
nhead_kv,
True,
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=PAGE_SIZE,
kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=max_q_len,
uni_seqlen_qo=max_q_len,
fast_mode=False,
max_split_per_batch=num_kv_splits,
intra_batch_mode=True,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
return {
"work_meta_data": work_metadata,
"work_indptr": work_indptr,
"work_info_set": work_info_set,
"reduce_indptr": reduce_indptr,
"reduce_final_map": reduce_final_map,
"reduce_partial_map": reduce_partial_map,
}
info = get_mla_metadata_info_v1(
batch_size, max_q_len, nhead, q_dtype, kv_dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=num_kv_splits,
intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
(work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
nhead // nhead_kv,
nhead_kv,
True,
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=PAGE_SIZE,
kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=max_q_len,
uni_seqlen_qo=max_q_len,
fast_mode=False,
max_split_per_batch=num_kv_splits,
intra_batch_mode=True,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
self._meta = (work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map)
self._meta_key = key
return {
"work_meta_data": work_metadata,
"work_indptr": work_indptr,
"work_info_set": work_info_set,
"reduce_indptr": reduce_indptr,
"reduce_final_map": reduce_final_map,
"reduce_partial_map": reduce_partial_map,
}
def get_dequant_bf16(self, total_kv: int, qk_head_dim: int) -> torch.Tensor:
"""
获取用于 MXFP4 反量化的 bf16 缓冲区。
当使用 MXFP4 格式的 KV 缓存时,需要先将 FP4 反量化为 bf16,
这个缓冲区用于存储反量化的中间结果。
参数:
total_kv: 所有 batch 中 KV token 的总数
qk_head_dim: QK 头维度 (576)
"""
key = (total_kv, qk_head_dim)
if self._dq_bf16_key != key or self._dq_bf16_buf is None:
self._dq_bf16_buf = torch.empty(
(total_kv, NUM_KV_HEADS, qk_head_dim),
dtype=torch.bfloat16, device="cuda",
)
self._dq_bf16_key = key
return self._dq_bf16_buf
# 创建全局单例缓存对象,在整个推理过程中复用
_cache = _BufferCache()
def _quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
动态逐张量 FP8 量化(dynamic per-tensor FP8 quantization)。
=== 量化原理 ===
FP8(8-bit 浮点数)只有 8 位来表示数值,范围和精度都有限。
为了最大化利用 FP8 的表示范围,我们需要:
1. 找到输入张量的最大绝对值 (amax)
2. 计算缩放因子 scale = amax / FP8_MAX
3. 将每个元素除以 scale,映射到 FP8 的完整范围
4. 裁剪到 FP8 可表示的范围,然后转换类型
反量化公式:original ≈ fp8_tensor * scale
参数:
tensor: 待量化的 bf16 张量
返回:
(fp8_tensor, scale):
- fp8_tensor: 量化后的 FP8 张量
- scale: 标量缩放因子(float32),形状为 (1,)
"""
# 找到整个张量的最大绝对值,clamp 确保不为零(避免除零)
amax = tensor.abs().amax().clamp(min=1e-12)
# 计算缩放因子:amax 映射到 FP8 的最大可表示值
scale = amax / FP8_MAX
# 量化:除以 scale → 裁剪到 FP8 范围 → 转换为 FP8 类型
fp8_tensor = (tensor / scale).clamp(min=-FP8_MAX, max=FP8_MAX).to(FP8_DTYPE)
return fp8_tensor, scale.to(torch.float32).reshape(1)
def _dequantize_mxfp4_inplace(
fp4_data: torch.Tensor,
scale_e8m0: torch.Tensor,
out_bf16: torch.Tensor,
) -> torch.Tensor:
"""
就地反量化 MXFP4 数据为 bf16。
=== MXFP4 格式说明 ===
MXFP4 (Microscaling FP4) 是一种块级量化格式:
- 每 32 个连续元素为一个块(block)
- 每个元素只有 4 bit(FP4 E2M1 格式:2位指数 + 1位尾数 + 1位符号)
- 每个块共享一个 E8M0 格式的缩放因子(8位纯指数,无尾数)
- 两个 FP4 值打包(pack)在一个字节(byte)中
因此对于 N 个元素:
- fp4_data 的形状为 (rows, N//2),每个字节包含 2 个 FP4 值
- scale_e8m0 的形状为 (rows, N//32),每 32 个元素共享 1 个缩放因子
=== 反量化过程 ===
1. 将打包的 FP4 数据解包为 float32(mxfp4_to_f32)
2. 将 E8M0 缩放因子转换为 float32(e8m0_to_f32)
3. 将数据按块 reshape,乘以对应的缩放因子
4. 将结果复制到 bf16 输出缓冲区
参数:
fp4_data: 打包的 FP4 数据,形状 (rows, N//2)
scale_e8m0: E8M0 块缩放因子,形状 (rows, num_blocks)
out_bf16: 预分配的 bf16 输出缓冲区,形状 (total_kv, 1, N)
"""
total_kv = fp4_data.shape[0]
N = out_bf16.shape[-1] # 原始维度 = 576
num_rows = total_kv
block_size = 32 # MXFP4 的块大小固定为 32
num_blocks = N // block_size # 576 / 32 = 18 个块
# 步骤 1: 将打包的 FP4 数据解包为 float32
# fp4_data 形状 (rows, N//2=288) → float_vals 形状 (rows, N=576)
fp4_data_2d = fp4_data.view(num_rows, N // 2)
float_vals = mxfp4_to_f32(fp4_data_2d)
# 步骤 2: 将 E8M0 缩放因子转换为 float32
# scale_e8m0 可能包含填充(padding),需要裁剪到实际需要的行数和块数
scale_f32 = e8m0_to_f32(scale_e8m0)
scale_f32 = scale_f32[:num_rows, :num_blocks]
# 步骤 3: 将数据按块 reshape,然后乘以对应的缩放因子
# float_vals: (rows, 576) → (rows, 18, 32)
# scale_f32: (rows, 18) → (rows, 18, 1) 通过 unsqueeze 广播
float_vals_blocked = float_vals.view(num_rows, num_blocks, block_size)
float_vals_blocked.mul_(scale_f32.unsqueeze(-1)) # 就地乘法,节省显存
# 步骤 4: 将结果复制到 bf16 输出缓冲区
# (rows, 18, 32) → (rows, 576) → (rows, 1, 576)
out_bf16.copy_(float_vals_blocked.view(num_rows, 1, N))
return out_bf16
def _get_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
"""
根据工作量动态计算 KV 分片数量(num_kv_splits)。
=== 什么是 KV 分片? ===
在 FlashAttention 等高效注意力算法中,KV 序列可以被分成多个片段(split),
每个 GPU 线程块独立处理一个片段,最后通过归约(reduction)合并结果。
分片数量的权衡:
- 分片太少:GPU 利用率低,性能差
- 分片太多:归约开销增大,且可能引入精度损失
=== 动态策略 ===
根据总工作量 total_work = batch_size × kv_seq_len 来决定:
- 工作量小(≤ 4096):不分片(1),因为数据量小,单线程块就够
- 工作量中等:逐步增加分片数(2, 4, 8, 16)
- 工作量大(> 262144):最多 32 个分片
注意:reference.py 中固定使用 32 个分片,而 submission.py 使用动态策略,
在工作量较小时可以减少不必要的归约开销。
"""
total_work = batch_size * kv_seq_len
if total_work <= 4096:
return 1
elif total_work <= 16384:
return 2
elif total_work <= 65536:
return 4
elif total_work <= 131072:
return 8
elif total_work <= 262144:
return 16
else:
return 32
def _run_a8w8(
q_fp8: torch.Tensor,
q_scale: torch.Tensor,
kv_fp8: torch.Tensor,
kv_scale: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict,
num_kv_splits: int,
) -> torch.Tensor:
"""
使用 A8W8(Activation-8bit, Weight-8bit)模式运行 MLA 解码注意力。
=== A8W8 模式 ===
"A8W8" 表示 Q(激活)和 KV(权重)都使用 8-bit 浮点数(FP8)。
这是 aiter 库中最快的注意力计算模式。
由于 FP8 的范围有限,需要提供缩放因子(scale)来恢复原始数值范围:
- q_scale: Q 的缩放因子,用于在计算中恢复 Q 的数值范围
- kv_scale: KV 的缩放因子,用于恢复 KV 的数值范围
=== 注意力计算流程 ===
1. 将 KV 缓存 reshape 为 4D 格式(total_kv, page_size, nhead_kv, dim)
2. 准备 KV 索引和元数据
3. 调用 mla_decode_fwd 执行核心注意力计算
4. 返回输出张量
参数:
q_fp8: FP8 量化的 Query,形状 (total_q, num_heads, qk_head_dim)
q_scale: Q 的缩放因子,标量 float32
kv_fp8: FP8 量化的 KV 缓存,形状 (total_kv, 1, 576)
kv_scale: KV 的缩放因子,标量 float32
qo_indptr: Query/Output 的索引指针,形状 (batch_size+1,)
类似 CSR 格式的行指针,qo_indptr[i]:qo_indptr[i+1] 是第 i 个 batch 的 Q
kv_indptr: KV 的索引指针,形状 (batch_size+1,)
kv_indptr[i]:kv_indptr[i+1] 是第 i 个 batch 的 KV
config: 包含注意力参数的配置字典
num_kv_splits: KV 分片数量
"""
# 从配置字典中提取参数
batch_size = config["batch_size"] # 批大小
nq = config["num_heads"] # Q 头数 = 16
nkv = config["num_kv_heads"] # KV 头数 = 1
dq = config["qk_head_dim"] # QK 维度 = 576
dv = config["v_head_dim"] # V 维度 = 512
q_seq_len = config["q_seq_len"] # Q 序列长度(解码时通常为 1)
total_kv_len = int(kv_indptr[-1].item()) # 所有 batch 的 KV token 总数
# 将 KV 缓存从 3D (total_kv, 1, 576) reshape 为 4D (total_kv, page_size, nhead_kv, 576)
# 这是 aiter mla_decode_fwd 要求的输入格式
kv_buffer_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, nkv, kv_fp8.shape[-1])
# 获取 KV 索引数组(从缓存中获取,避免重复分配)
kv_indices = _cache.get_kv_indices(total_kv_len)
# 计算每个 batch 的最后一个页的有效长度
# 由于 PAGE_SIZE=1,每个 batch 的 kv_last_page_len 就是该 batch 的 KV 长度
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
max_q_len = q_seq_len
# 获取或创建持久模式的元数据(从缓存中获取)
meta = _cache.get_metadata(
batch_size, max_q_len, nq, nkv,
q_fp8.dtype, kv_fp8.dtype,
qo_indptr, kv_indptr, kv_last_page_len,
num_kv_splits,
)
# 获取输出缓冲区(从缓存中获取,避免重复分配)
o = _cache.get_output(q_fp8.shape[0], nq, dv)
# ====================================================================
# 核心调用:aiter MLA 解码前向内核
# ====================================================================
#
# mla_decode_fwd 是 AMD aiter 库提供的高性能 MLA 解码注意力内核。
# 它在 GPU 上执行以下计算:
#
# O_i = softmax(Q_i @ K^T / sqrt(d)) @ V
#
# 其中:
# Q_i: 第 i 个 query token 的吸收表示 (16 heads × 576 dim)
# K: 所有 KV 缓存的键部分 (576 dim),前 512 维 + 64 维 RoPE
# V: 所有 KV 缓存的值部分 (前 512 维)
# O_i: 输出 (16 heads × 512 dim)
#
mla_decode_fwd(
q_fp8.view(-1, nq, dq),
kv_buffer_4d,
o,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_len,
max_q_len,
page_size=PAGE_SIZE,
nhead_kv=nkv,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=num_kv_splits,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
**meta,
)
return o
def custom_kernel(data: input_t) -> output_t:
"""
主入口函数:接收输入数据,执行 MLA 解码注意力计算。
=== 调用流程 ===
1. 解包输入数据(Q, KV 缓存, 索引指针, 配置)
2. 根据工作量计算 KV 分片数量
3. 将 Q 从 bf16 量化为 FP8
4. 从 KV 数据中提取 FP8 格式的 KV 缓存(已预量化)
5. 调用 A8W8 内核执行注意力计算
=== 输入数据结构 ===
data = (q, kv_data, qo_indptr, kv_indptr, config)
- q: (total_q, 16, 576) bf16 — 吸收后的 Query
- kv_data: 字典,包含三种精度的 KV 缓存:
- "bf16": (total_kv, 1, 576) bf16 — 最高精度
- "fp8": (kv_fp8, scale) — FP8 量化格式(本实现使用这个)
- "mxfp4": (fp4_data, scale) — MXFP4 量化格式(备选)
- qo_indptr: (batch_size+1,) int32 — Q/Output 的索引指针
- kv_indptr: (batch_size+1,) int32 — KV 的索引指针
- config: 配置字典(batch_size, 维度信息等)
=== 输出 ===
output: (total_q, 16, 512) bf16 — 注意力计算结果
"""
q, kv_data, qo_indptr, kv_indptr, config = data
num_kv_splits = 32
q_fp8, q_scale = _quantize_fp8(q)
kv_buffer_fp8, kv_scale = kv_data["fp8"]
return _run_a8w8(
q_fp8, q_scale,
kv_buffer_fp8, kv_scale,
qo_indptr, kv_indptr,
config, num_kv_splits,
)
scrolls · 572 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON