DeepSeek 机器学习系统工程师刘胜与的一篇长文引发了广泛讨论。比情绪表达更值得工程师关注的,是他的具体工作:独立实现并交付 DeepSeek V4.1 的主 Attention 算子,其中采用了 head_dim = 512 的 MQA Attention。这个信息把一个常被抽象成公式的模块,拉回到真实的系统工程现场:算子不仅要算对,还要在指定硬件、精度和推理负载下稳定地跑快。
现有摘要没有披露 V4.1 的完整模型结构、内核实现或性能数据,因此不能据此推断它使用了哪套 CUDA 技巧。我们可以做的,是从公开的算子形态出发,分析这项工作的难点,并搭建一个可复用的基准测试框架。
head_dim = 512 为什么值得注意
标准多头注意力中,每个 Query 头通常拥有对应的 Key 和 Value 头。MQA,也就是 Multi-Query Attention,则让多个 Query 头共享较少的 KV 头;最典型的情况是所有 Query 头共享一组 K 和 V。
假设张量形状如下:
Q: [batch, query_heads, sequence, head_dim]
K: [batch, 1, sequence, head_dim]
V: [batch, 1, sequence, head_dim]
MQA 的直接收益是缩小 KV Cache。若 Query 有 16 个头,而 KV 只有 1 个头,单层 KV Cache 的元素数量可以降到对应多头结构的约 1/16。对长上下文和并发推理而言,这会直接影响显存容量、内存带宽以及可容纳的请求数量。
但 head_dim = 512 也给内核实现带来压力。单个注意力头需要处理更宽的向量,点积和归一化涉及更多数据;寄存器占用、线程块划分、共享内存使用和中间结果保存都会变得更敏感。一个在较小 head dimension 上表现良好的 kernel,不能假定放大到 512 后仍然高效。
真正的交付不止是写出 CUDA Kernel
一个主 Attention 算子进入模型主路径,至少要同时守住四条边界:
- 数值正确性:不同序列长度、batch、mask 和精度下,结果必须与可信参考实现保持在误差范围内。
- 性能稳定性:不能只在某个整齐形状上快,还要覆盖预填充、逐 token 解码和不同并发度。
- 资源约束:寄存器过多可能降低 occupancy,中间张量物化则可能抵消 MQA 节省显存的意义。
- 工程可维护性:算子需要接入构建、测试、模型调度和回退路径,并能在硬件或编译器升级后重新验证。
因此,“独立实现并交付”比单次 benchmark 跑出高分更重。它意味着实现者需要把算法、GPU 执行模型和生产系统连在一起,并对最终行为负责。
可以这样实践:建立一个 MQA 基准测试基线
下面的脚本不是 V4.1 算子的复现,也不代表其内部实现。它提供一个可以直接运行的 PyTorch 基线,用来比较两种 MQA 执行方式:显式复制 KV 头,以及使用 PyTorch SDPA 的 GQA/MQA 支持。
运行前安装 PyTorch 2.x;有 CUDA 时脚本会自动使用 FP16 和 GPU,没有 CUDA 时会缩短序列并使用 CPU。
# bench_mqa.py
import time
import torch
import torch.nn.functional as F
def sync(device):
if device.type == "cuda":
torch.cuda.synchronize()
def measure(fn, device, warmup=5, runs=20):
for _ in range(warmup):
fn()
sync(device)
start = time.perf_counter()
for _ in range(runs):
fn()
sync(device)
return (time.perf_counter() - start) * 1000 / runs
def main():
torch.manual_seed(0)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dtype = torch.float16 if device.type == "cuda" else torch.float32
batch = 1
query_heads = 16
kv_heads = 1
sequence = 512 if device.type == "cuda" else 128
head_dim = 512
q = torch.randn(
batch, query_heads, sequence, head_dim,
device=device, dtype=dtype,
)
k = torch.randn(
batch, kv_heads, sequence, head_dim,
device=device, dtype=dtype,
)
v = torch.randn_like(k)
# 参考路径:物化复制后的 K/V,因此会额外占用显存。
k_repeated = k.repeat_interleave(query_heads // kv_heads, dim=1)
v_repeated = v.repeat_interleave(query_heads // kv_heads, dim=1)
def repeated_kv():
return F.scaled_dot_product_attention(
q, k_repeated, v_repeated, is_causal=True
)
def native_mqa():
return F.scaled_dot_product_attention(
q, k, v, is_causal=True, enable_gqa=True
)
reference = repeated_kv()
try:
candidate = native_mqa()
except TypeError as exc:
raise SystemExit(
"当前 PyTorch 不支持 enable_gqa,请升级 PyTorch 2.x。"
) from exc
torch.testing.assert_close(
candidate,
reference,
rtol=2e-2 if dtype == torch.float16 else 1e-4,
atol=2e-2 if dtype == torch.float16 else 1e-5,
)
repeated_ms = measure(repeated_kv, device)
native_ms = measure(native_mqa, device)
print(f"device={device}, dtype={dtype}, shape={tuple(q.shape)}")
print(f"repeated KV: {repeated_ms:.3f} ms")
print(f"native MQA : {native_ms:.3f} ms")
print(f"speed ratio: {repeated_ms / native_ms:.2f}x")
if __name__ == "__main__":
main()
执行命令:
python bench_mqa.py
若要进一步定位 GPU 时间和 kernel 调度,可以使用 Nsight Systems:
nsys profile --trace=cuda,nvtx,osrt -o mqa_profile python bench_mqa.py
这个实验不能证明某个生产算子更快,但能建立三个重要习惯:先准备可信参考实现,再检查数值误差,最后才比较延迟。测试时还应增加不同 batch、上下文长度、Query 头数以及解码阶段 q_len = 1 的组合,避免只优化一个漂亮的形状。
从算子到工程师的长期选择
这次讨论之所以超出技术圈,部分原因是个人表达与高强度系统工作形成了反差。顶层模型能力常被迅速看见,底层算子的价值却藏在延迟、吞吐、显存占用和线上稳定性里。一次有效的优化往往很快成为下一版系统的默认起点,工程师随即面对新的瓶颈。
对准备投入类似工作的团队,可以用一份简短清单约束项目:明确目标硬件和真实负载;保留可读的参考实现;建立数值、性能和显存三类回归测试;记录编译器与框架版本;为不支持的形状设计回退路径。至于个人是否继续留在这场加速中,则不只是技术判断,还涉及成长空间、工作节奏和成果归属。持续优化有价值,但让经验能够沉淀、复现和传递,同样是工程成果的一部分。