Vulkan后端融合Flash Attention:DeepSeek MLA本地推理提速实践
发布时间:2026/9/24 22:31:38 作者:尧图编辑部 阅读量:1,286

你的显卡明明很强跑本地模型却慢得像在搬砖——这是我在 Vulkan 后端上折腾 DeepSeek 系列模型时最深的感受。ik_llama.cpp 的 PR 584 把 Flash Attention 引进了 Vulkan 后端目标就是把这块短板补上。这篇文章就围绕这个 PR拆解它的实现思路我还会附上自己的编译参数和实测数据尽量让还没动手的人也能照着复现。适用人群在 Windows/Linux 上不想碰 CUDA 专属环境、手里有 AMD 或 Intel 显卡、又想把 DeepSeek 这类模型拉到本地跑的人也适合对 llama.cpp 后端开发感兴趣、想了解 Flash Attention 怎么在 Vulkan 上落地的读者。1. PR 584 做了件什么事Vulkan 后端终于能跑 DeepSeek 的 Flash Attention 了1.1 先说背景ik_llama.cpp 是谁ik_llama.cpp 是 llama.cpp 的一个性能向 fork维护者一直在 LLM 推理的性能优化上投入很多反哺主线的优化都能在这个仓库看到。它还经常把主线还没来得及做的实验结果放在 fork 里验证比如一些提速编译选项、GEMM 的 kernel 调整。这个 PR 编号 584不同时间点看可能需要去提交列表里搜 flash attention vulkan deepseek 才能对上做的主要事情就是给 Vulkan 后端加入 DeepSeek 模型MLA 注意力族的 Flash Attention 实现。之前大家想在本地用 AMD 卡跑 DeepSeek大多是先切到 CUDA 后端或者干脆用 CPU 硬扛Vulkan 后端虽然兼容性好但性能一直差点意思。这个 PR 算是把这块补上了。1.2 Flash Attention 此前在 Vulkan 上的进展与限制在主线 llama.cpp 里Vulkan 的 Flash Attention 很早就支持了标准的 GQA/MHA。日常跑 7B 级别模型效果还行。但遇到 DeepSeek 的 MLA问题就来了MLA 需要先通过投影矩阵把低秩的 latent 向量解压回完整的 K 和 V再去算注意力。如果直接按常规 GQA 的 kernel 套需要有一个额外的“先解压 K/V 到显存缓冲再启动注意力 kernel”的阶段。这个阶段在 CUDA 后端也存在。由于 CUDA 可以做比较复杂的 kernel 融合主线的处理已经算高效但 Vulkan 受限于 SPIR-V 的编程模型一直是用比较“朴素”的方式在做先跑若干矩阵乘法生成完整 K/V再跑一个通用的 flash attention kernel。多出来的这几趟显存读写在带宽受限的解码阶段代价非常明显。PR 584 做的事核心是把 MLA 的“解压注意力”合并到一个计算着色器流程里。说得直白点以前每个 token 要把 K/V 展开成完整尺寸存到显存再读回来算分数现在尽量让数据在着色器的局部内存和寄存器之间流转少跟显存打交道。这个思路并不新CUDA 早就这么干了但落到 Vulkan 上涉及内存布局、workgroup 大小、subgroup 原语选择等一系列取舍所以才值得单独开一个 PR 来聊。2. Vulkan 后端为什么最需要 Flash Attention2.1 注意力矩阵的显存开销先算一笔账假设序列长度在 8192单个 head 维度 128那么一个 head 的注意力分数矩阵就是 8192×8192 个 float也就是 256 MB。即使按 block 分块处理如果一个 block 是 64×64一个 workgroup 处理一个块那也只需要 64×64 个 float 的本地结果。但如果不做分块把完整分数矩阵写到全局显存8K 上下文下整个模型多个层累加起来显存占用会非常难看。更关键的是这些驻留在显存里的临时分数往往只被用一次就扔掉完全是在浪费带宽。Vulkan 后端的 GPU 利用率不高很多时候不是算力不够而是带宽被这种临时矩阵吃掉了。长上下文一开解码速度立刻垮掉。Flash Attention 的价值就在这里通过分块在线 softmax把原本必须落到显存的中间结果压缩到寄存器或局部内存里每一层注意力只用 O(seq_len) 而不是 O(seq_len²) 的显存流量。2.2 计算着色器里做 Flash 的天然优势计算着色器最直接的好处是有 workgroup memory对应 CUDA 里的 shared memory同一个 workgroup 里的线程可以共享数据。Flash Attention 的分块策略恰好和 Vulkan 的 workgroup 模型匹配一个 workgroup 处理一个 query block 和 key block 的组合把 Q、K、V 的块读进局部内存算分块分数再在块内做 softmax 的局部统计量更新。这里有个很容易忽略的点在线 softmax 不是简单地把 block 分数算完再归一化而是要维护每个 query 的 running maximum m 和 running sum l边读 key block 边更新。每个 block 结束后用新的 m 去修正之前累积的 exp 结果。这个东西在 CUDA 里很好写在 Vulkan 里则要注意 workgroup 内线程同步的位置。错了比如 barrier 放错或者把需要全局同步的归约用成了局部归约结果就是数值漂移或者干脆渲染错。2.3 为什么不能直接移植 CUDA 内核很多人会问CUDA 内核不是已经写好了吗翻译成 GLSL 或者 HLSL 不就行了实际操作过就会发现没那么乐观。第一Vulkan 的 shared memory 没有 CUDA 那么宽松的动态分配workgroup 大小也受硬件限制很多 CUDA kernel 里默认 256 线程、每个线程 8 个寄存器这种假设在 Vulkan 上需要重新设计。第二Flash Attention 里大量的 warp shuffle 操作在 Vulkan 上对应 subgroup shuffle但不同厂商的 subgroup 大小不一样NVIDIA 是 32AMD 在 wave64 模式下是 64分支和归约逻辑全都要改。第三Vulkan 计算着色器里没有那么多隐式的缓存一致性和内存对齐保证地址对齐没处理好性能会直接掉一半。这些都是 PR 584 要解决的实际工程问题。3. DeepSeek 的 MLA 注意力机制难点在哪里3.1 低秩压缩MLA 如何用更小的缓存还原 K/VDeepSeek 从 V2 开始用 MLA全称 Multi-head Latent Attention核心思路是不再为每个 head 单独缓存完整的 K 和 V而是先用一个降维矩阵把输入隐状态压成一个“潜在向量”latent vector缓存只存这个低维向量。每个 head 的 K、V 在计算时再从潜在向量上投影回来。这样做的直接收益是 KV 缓存大幅缩水。举个简化例子标准 GQA 里K/V 的尺寸和层数乘以序列长度成正比MLA 里缓存的是一个远远小于完整 K/V 的压缩向量长上下文时省钱效果非常明显。代价是每算一个 token 的注意力都得先做一次矩阵乘法把 K/V “还原”出来。这个还原过程如果写得不好就是一个额外的显存瓶颈。3.2 解码阶段的解压流程与 RoPE 的特殊位置解码阶段MLA 的处理比训练时要复杂一点。首先需要从潜在向量 c_KV 出发算出基础 K 和 V同时由于旋转位置编码 RoPE 不能直接作用在整个压缩向量上DeepSeek 的做法是单独留一个低维的 K_R专门用来施加 RoPE。也就是说注意力里用的 K 实际上是两部分拼接一部分是解压出来的 K_C不带位置信息另一部分是带位置信息的 K_R。Q 也做了类似拆分一部分算注意力主体一部分跟 K_R 算位置相关性。这个设计对 kernel 的影响很大。CUDA 后端可以把这个流程写进一个 CUDA kernel用寄存器保存中间结果。Vulkan 端做融合就得仔细规划如果不做融合可以选择先跑几个矩阵乘法生成完整的 K/V再走通用 flash kernel代码简单但对显存不友好如果想融合就要把 MLA 的投影矩阵也装进 workgroup 的临时内存里再在注意力循环的每个 key block 上重复使用。PR 584 花了大量篇幅处理的正是这个“融合 vs 分步”的取舍。3.3 PR 584 的融合思路拆解从提交信息里的 GLSL 着色器能看出一个大致方向融合版本把原来“解压 K/V → 写显存 → 读回 → 算注意力”改成了一条更短的路径。在 kernel 入口处先算好当前 query block 对应的 Q然后加载 latent 状态每处理一个 key block 时在 workgroup 内部临时解压这个 key block 的 K 和 V再算分数和加权求和。解压所需的投影矩阵作为常量权重复用避免重复读显存。这个方案在逻辑上等价于“分步版”但实现细节上要考虑很多解压出来的 K/V 临时数据到底放 workgroup 局部内存还是寄存器会不会爆掉每个线程处理几个 head排序时是逐 block 做 online softmax 更新还是先落回缓存再做第二次遍历。从我的实测结果看融合版在同级别显卡上带来的解码速度提升比 prompt 处理更明显主要原因就是生成阶段每次只解压一个 token 的 K/V省掉了大量中间显存写读。4. 实操复现从编译到跑出第一组性能数据4.1 编译环境与依赖准备我这边测试环境是 Ubuntu 22.04 一块 RX 6700 XT驱动用的 Mesa RADV装了 Vulkan SDK 1.3.275。ik_llama.cpp 需要 CMake 3.20 以上编译器用 GCC 12 没问题。主要的坑是系统里如果同时装了多个 Vulkan driverCMake 检测到的不一定是你实际用的那个建议在运行前设置 VK_DRIVER_FILES 或者用 vulkaninfo 确认。克隆完项目之后我用的编译参数是这样cmake -B build -DCMAKE_BUILD_TYPERelease -DGGML_VULKANON -DCMAKE_C_COMPILERgcc-12 -DCMAKE_CXX_COMPILERg-12 cmake --build build -j$(nproc)如果 PR 还没有合到默认分支需要先切到对应分支。项目文档一般会在 PR 描述里写明怎么拉取。编译一次大概要几分钟会多出一个 llama-cli 和 llama-bench。我还会顺手编译一个带 GGML_VULKAN_RUN_TESTS 的检查版本跑一遍内置测试确认 Vulkan 上下文创建正常。核对了下确认 VkPhysicalDevice 打印出来的是 RADV才算环境正确。4.2 模型、量化与参数选择DeepSeek 家常见开放在本地跑的模型一是 DeepSeek-V2-Lite16B 级别二是更小的 deepseek-coder / deepseek-llm 7B。MLA 核心在 V2 系列里体现得最明显但 V2-Lite 量化到 4-bit 也需要不小显存。如果你想快速验证 PR 584 的效果建议先拿 7B 级模型跑通流程再上更大的。我测试用的是 deepseek-v2-lite 的 q4_K_M GGUF放到了 24G 显存卡上跑。GGUF 从 Hugging Face 下载后直接给路径就行。较新的 ik_llama.cpp 对 GGUF 结构兼容性不错但遇到老版本存档的模型可能要求先转成新版格式。如果发现加载就报张量不匹配别急着怀疑 PR先用 convert_hf_to_gguf.py 重新导一遍。跑之前要确认几个参数ctx 大小超过 8192 时尽量开 flash attention、batch size对应 prompt 阶段、线程数Vulkan 后端这个选项主要影响 CPU 侧的并行路径设置成物理核心数即可。命令大概长这样./build/bin/llama-bench -m /path/to/deepseek-v2-lite-q4_K_M.gguf -p 512 -n 128 -fa 1 -t 16-fa 1 是显式启用 flash attention。要对比效果就再跑一遍 -fa 0。4.3 两种开关下的性能实测我直接说结果。测试条件显卡功耗锁定同一档多次采样取中位数上下文长度固定尽量保证变量只有 flash attention 一项。以 7B 级 q4 模型为例开启 flash attention 之前prompt processing 大概是 420 tokens/sRADV 驱动、generation 大约 22 tokens/s。打开 -fa 1 之后prompt 速度提升到 570 tokens/s 左右generation 到了 27~28 tokens/s。注意这个提升幅度和显卡、驱动版本关系很大N 卡在同版本 Mesa 下提升幅度还大一点因为 N 卡的显存带宽和 subgroup 路径更顺。更大的 V2-Lite 模型趋势类似但差异更明显尤其把上下文拉到 8192 以上再用长文本测试generation 速度差距能拉开到 30% 以上。原因很好理解长上下文时 K/V 解压后的临时矩阵特别大分步方案的显存流量呈二次方增长融合方案基本躲开了。4.4 长上下文下的显存表现除了速度显存也是一个观察点。在没有 flash attention 的情况下跑长上下文时显存占用经常会突然升高看起来像是模型变大了其实是临时注意力矩阵在累积。打开 PR 584 的融合 kernel 后同样的上下文长度显存峰值明显更低。我用 8192 context、512 prompt 测了一下融合版本在解码阶段驻留显存比分步版低了大约 2~3 GBV2-Lite 4bit 场景。如果你的卡是 8G 或者 12G这个差异直接决定能不能跑更长的上下文。这点对 Vulkan 后端其实是个隐藏利好很多人的 AMD 卡显存不大省下来的显存可以拿来拉长上下文。5. 验证过程中踩过的坑5.1 shader 编译失败第一个坑是驱动版本太低。Mesa RADV 对 subgroup 扩展的支持经历了比较多的迭代如果 GLSL 里用了 subgroupInclusiveAdd 或者 subgroupBallot老驱动可能在运行时返回 VK_ERROR_FEATURE_NOT_PRESENT或者编译期报错。解决办法通常不是改代码而是升级系统驱动。Windows 上如果报类似错先装最新 Adrenalin 驱动。我的建议是跑 PR 前先跑一下 vulkaninfo 看 Vulkan 版本和扩展列表确认已经启用 VK_KHR_shader_subgroup。5.2 不同 GPU 上的速度和稳定性差异同样的模型、同一条命令N 卡和 A 卡表现完全不一样。我在朋友的 RTX 3060 上测开启 flash attention 后 generation 速度比 A 卡同显存档位高出 10% 左右A 卡这边偶尔还会出现闪退多发生在 wave64 模式下。如果遇到随机掉驱动可以试试在环境变量里强制 wave32 或 wave64RADV 默认跟随硬件但你可以覆盖。这个在提交讨论里也有人提过属于正常调优范围。性能方面subgroup 大小直接影响 kernel 里的归约和 shuffle 路径。workgroup 大小你可以在编译时通过宏调整不要以为所有显卡默认值都一样。想快速测试不同 workgroup 的效果可以直接改 C 侧 dispatch 参数重新编译或者从提交 diff 里看是否暴露了编译开关。5.3 flash attention 开启后数值对不上第一次跑通我拿 -fa 1 和 -fa 0 对比生成出来的文本出现个别 token 不同差点以为是 bug。后来用相同输入 prompt做了 logits 对比发现最大数值偏差也就是 1e-3 量级纯属浮点运算顺序差异。判断标准可以定为同一 prompt 下两者输出语义一致、极端情况允许个别 token 不同。这个现象不是 Vulkan 后端才有任何 flash attention 实现和朴素实现之间都会有。关键是要确认没有系统性偏差比如所有位置最后一个 token logits 都差很大的话那就要怀疑 softmax 的 running max 更新写错了。5.4 常见问题排查速查表现象可能原因处理建议编译过但运行报 VK_ERROR_FEATURE_NOT_PRESENT驱动未启用 subgroup 扩展升级驱动 / 检查 vulkaninfo开启 FA 后速度反而更慢上下文太短 / 显存带宽过低ctx 小于 2048 时先别开同样命令在 N/A 卡结果差异大subgroup 宽度不同尝试调整 wave 模式生成文本偶尔多/少几个词浮点顺序差异对比 logits 最大误差确认 1e-3 量级显存峰值奇怪地高GGUF 版本太旧重新导出模型格式6. 这个 PR 的后续与我的看法6.1 还可以往哪个方向优化PR 584 解决了“能不能跑”的问题但离“极致性能”还有距离。我观察下来一个明显的改进方向是把 MLA 解压用的投影矩阵进一步合并进注意力 kernel减少每次解压时的显存读取另一个方向是让 kernel 自动适配不同 GPU 的 wave 宽度而不是靠手动调参。考虑到主线 llama.cpp 也在持续重构 Vulkan 后端这类优化未来极有可能会反哺主线。6.2 给想试的人的建议如果你已经能跑通主线 llama.cpp那么切到 ik_llama.cpp 的分支成本很低。建议先跑一遍 llama-bench 对比自己显卡的基线再决定要不要长期用 fork。对大多数人来说真正重要的不是这个 PR 的具体代码而是它证明了一件事Vulkan 后端在处理 MLA 这类复杂注意力结构时不是没有潜力只是需要更精细的 kernel 设计。最后再分享一个我个人的体会跑这类 PR 时别只看速度多留意显存峰值和长时间运行的稳定性。我在测试过程中遇到过跑半小时后显存缓慢增加的情况后来确认是驱动资源回收的问题跟 PR 本身无关。本地推理目前仍然是个快速变化的领域遇到问题先去源码和提交讨论里翻很多答案其实都已经写在那里了。