DeepSeek开源周第一天:FlashMLA 深度分析
DeepSeek FlashMLA 深度分析

今天早上9点34,deepseek在X上发布了开源周第一天的项目FlashMLA,这篇就带来FlashMLA的深度分析。

FlashMLA项目非常受欢迎,目前代码已经有6.8k Star
简单介绍MLA
MLA(Multi-Head Latent Attention) 是deepseek 在论文 DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model 中提出来的针对多头注意力(Multi-Head Attention, MHA)的优化方法。
在Transformer模型中,MHA是最消耗算力的模块之一,如果想要大规模场景下保持高效率,就需要进一步优化。

MLA算是MHA的一个变种,在实现上继承了FlashAttention的一些东西,DeepSeek-V2论文中主要是对比MHA,GQA,MQA,优化效果如下图:


在一些推理框架中,也实现了MLA。如下图所示,SGLang加入 MLA之后,吞吐量有2-3倍的增加。

FlashMLA 使用
环境要求是:
-
Hopper GPUs
-
最低CUDA 12.3
-
最低 PyTorch 2.0
安装:
git clone https://github.com/deepseek-ai/FlashMLA
python setup.py install
使用方法:

性能:
仓库提供了Benchmark文件可以直接跑。官方列出在 H800 SXM5 上,使用 CUDA 12.6,在内存限制配置下实现高达 3000 GB/s 的速度,在计算限制配置下实现 580 TFLOPS 的性能。
代码分析

FlashMLA的代码不多,依赖也很少。
主要的优化手段如下:
计算分块和调度优化
templatestruct Flash_fwd_kernel_traits_mla {
// 固定 block 大小为 64x64
static constexpr int kBlockM = kBlockM_;
static constexpr int kBlockN = kBlockN_;
// 每个 block 使用 8 个 warp 并行
static constexpr int kNWarps = kNWarps_;
static constexpr int kNThreads = kNWarps * 32;
// 共享内存优化
static constexpr int kBlockKSmem = kHeadDim % 64 == 0 ? 64 : 32;};
关键点:
- 通过分块(块大小64)分页KV缓存和多warp并行提高计算效率
内存访问优化
struct Flash_fwd_mla_params {
using index_t = int64_t;
int b, seqlen_q, d, d_v;
int h, h_h_k_ratio, ngroups;
bool is_causal;
float scale_softmax, scale_softmax_log2;
int *__restrict__ cu_seqlens_k;
void *__restrict__ q_ptr;
void *__restrict__ k_ptr;
void *__restrict__ v_ptr;
void *__restrict__ o_ptr;
void *__restrict__ softmax_lse_ptr;
index_t q_batch_stride;
index_t k_batch_stride;
index_t v_batch_stride;
index_t o_batch_stride;
index_t q_row_stride;
index_t k_row_stride;
index_t v_row_stride;
index_t o_row_stride;
index_t q_head_stride;
index_t k_head_stride;
index_t v_head_stride;
index_t o_head_stride;
int *__restrict__ block_table;
index_t block_table_batch_stride;
int page_block_size;
int *__restrict__ tile_scheduler_metadata_ptr;
int num_sm_parts;
int *__restrict__ num_splits_ptr;
void *__restrict__ softmax_lseaccum_ptr;
void *__restrict__ oaccum_ptr;};
关键点:
-
使用分页 KV 缓存(block_table, page_block_size)
-
优化的内存布局和访问步长(stride)
-
使用 tile_scheduler_metadata 进行调度
Softmax 计算优化
for (int mi = 0; mi (tensor); ++mi) {
MaxOp max_op;
max(mi) = zero_init ? tensor(mi, 0) : max_op(max(mi), tensor(mi, 0));
#pragma unroll
for (int ni = 1; ni (tensor); ni++) {
max(mi) = max_op(max(mi), tensor(mi, ni));
}
max(mi) = Allreduce::run(max(mi), max_op);
// If max is -inf, then all elements must have been -inf (possibly due to masking).
// We don't want (-inf - (-inf)) since that would give NaN.
const float max_scaled = max(mi) == -INFINITY ? 0.f : max(mi) * scale;
sum(mi) = 0;
#pragma unroll
for (int ni = 0; ni (tensor); ++ni)
{
// Instead of computing exp(x - max), we compute exp2(x * log_2(e) -
// max * log_2(e)) This allows the compiler to use the ffma
// instruction instead of fadd and fmul separately.
tensor(mi, ni) = exp2f(tensor(mi, ni) * scale - max_scaled);
sum(mi) += tensor(mi, ni);
}
SumOp sum_op;
sum(mi) = Allreduce::run(sum(mi), sum_op);}
关键点:
-
使用 log2/exp2 替代 log/exp
-
使用 FFMA 指令优化
-
通过 warp-level 规约优化求和
双缓冲优化
struct SharedStorageMLA {
union {
struct {
// 双缓冲 K 矩阵
cute::array_aligned> smem_q;
cute::array_aligned * 2> smem_k;
// Double buffer
cute::array_aligned> smem_p;
cute::array_aligned> smem_scale;
};
};};
关键点:
- 通过双缓冲隐藏内存延迟,提高硬件利用率
FlashMLA总的来说是算是定制版的flash attention,目前适用场景:
-
环境需要是CUDA 11+、SM90+ Hopper架构
-
需要对 BF16(Q=576, V=512)的多头注意力做推理或训练
-
对大序列场景,需结合 split-K 方案提升吞吐
昨天我猜的推理优化,应该算猜对?我们来复习下:开源周来袭:DeepSeek有哪些神秘技术即将公开?

如上图,目前官方优化的方法还有不少,期待明天的项目,感觉还是infra,难度会是MTP? 关注我,我们一起猛学。
链接:
-
DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model https://arxiv.org/abs/2405.04434
-
sglang-v0-3 https://lmsys.org/blog/2024-09-04-sglang-v0-3/