这几年大模型把上下文长度越卷越长128K 甚至 1M 的序列都开始走向工程落地。模型参数固定下来之后真正卡脖子的往往是注意力Attention那部分序列一长QK^T 产生的中间矩阵动辄几百 MB读写开销直接把推理和训练的速度拖垮。FlashAttention 之所以被当成救命稻草是因为它在算法层面把访存复杂度从 O(N²) 压回线性省掉了完整注意力矩阵的落盘。但算法省下来的 IO最终还得靠硬件接住。我最近做的一个方向就是把 FlashAttention 的算子落到一颗“寄存器级计算 3D 堆叠存储”的 NPU 上让注意力计算尽量在芯片内部的小存储里就地完成不往外挪数据。这篇文章把我这段时间的设计思路、踩过的坑和验证方法整理出来给同样在做 AI 加速器、算子库或 LLM 推理优化的朋友做参考。1. 三个关键词背后到底在解决什么问题“寄存器级 3D 堆叠 NPU 加速 FlashAttention”这句话拆开看其实是在做一件事把大模型注意力中最难伺候的访存密集环节压到最靠近计算单元的存储层级上。要理解为什么这么设计得先把三个关键词各自的痛点说清楚。1.1 FlashAttention省了 HBM 的流量但没省掉 SRAM 的压力FlashAttention 的核心思路不复杂把 Q、K、V 切成小块在片上分别计算局部注意力分数同时维护一个逐步更新的 softmax 统计量最后输出 O。这样做最大的好处是不用把完整的分数矩阵 S 写回显存也不用再从显存里把它读回来做 softmax理论上每个 block 的数据只被读一次。但严格来说FlashAttention 只是把“大 HBM 流量”换成了“片上 SRAM 读写”。我在 GPU 上调算子时发现当 sequence length 超过一定规模后S 矩阵的 tile 会在 SRAM 和寄存器之间反复搬动搬来搬去照样占时间。换到 NPU 上也是一样如果只做到 L2 cache 这一级算力再高也会被数据搬运卡住。真正能做到“分数矩阵根本不离开寄存器”的场合才配得上叫寄存器级 FlashAttention。1.2 NPU 的矩阵引擎天生比通用处理器适合这类算子通用 CPU 和 GPU 的强项是通用性但代价是调度开销和记忆体等待。NPU 更像一条刚性流水线控制逻辑少绝大多数晶体管都放在乘加单元和片上存储上。用 NPU 做 FlashAttention核心不是要比谁的 ALU 跑得快而是要比谁能在相同功耗下让数据流最顺畅地通过计算阵列。当然NPU 也不是万能的。它通常只擅长固定形状的稠密矩阵乘而注意力里既有 GEMM通用矩阵乘又有 row-wise 的 softmax、mask 这类规约操作。这就需要在硬件设计阶段把 PE 阵列的互连结构、寄存器文件布局和算法分块方式一起考虑而不是等芯片回来了再靠算子库硬凑。1.3 寄存器级 3D 堆叠把“最后一公里”的搬运也省掉我把这个芯片方案理解成两段式优化。第一段是用 3D 堆叠把存储和逻辑层放到一起缩短长距离片间传输第二段是在计算阵列内部使用寄存器文件作为数据驻留地让注意力分数的更新、softmax 的累计、输出累加这三件事都在 PE 附近的寄存器里完成。打个比方普通方案是原料从大仓库发到工厂门口HBM卸货到中转仓SRAM再送到工人手里寄存器。3D 堆叠把仓库直接搬到了工厂隔壁而寄存器级设计让工人拿了料就不撒手直到成品做完才还回去。省掉的不只是几次搬运的耗时还有搬运过程中需要做的同步、仲裁和流水线停顿。1.4 这个方案适合谁以及适合什么场景如果你只做短序列的 CV 模型比如 512 分辨率的检测任务FlashAttention 的收益其实不大寄存器级方案显得大材小用。但如果是长文本预训练、长上下文推理、多轮 RAG、端侧多模态理解这类序列动辄几万 token 的场景注意力计算在总耗时里占比能到 60% 以上这时候把注意力算子的访存优化做透整体收益非常明显。2. 寄存器级 3D 堆叠 NPU 的架构选型逻辑这部分说的是硬件本身的取舍。一开始团队也讨论过用 GPU HBM 的方案以及用 Chiplet 拼接的方案最后都否了。2.1 3D 堆叠的层次划分逻辑层、SRAM 层与 DRAM 层3D 堆叠 NPU 最常见的做法是分成三层。最底层放计算逻辑和部分 SRAM中间层放堆叠的 SRAM 或者小型 DRAM最顶层再通过高密度硅通孔TSV或者混合键合的方式堆叠多层 DRAM。这样做的好处是逻辑层和存储层之间的物理距离可以缩短到几十微米能同时获得高带宽和低延迟。我在设计里把 SRAM 层放在逻辑层正上方容量在 4MB 左右分成多块 bank这样 PE 阵列在访问 K、V block 时不会因为 bank 冲突而坐等。DRAM 层用 8 层堆叠的方式带宽做到比传统片外 DRAM 高一个数量级。为了不引入过于复杂的制造流程我并没有直接用混合键合而是保留了 TSV 作为第一版方案TSV 在成本控制和可测试性上都成熟一些。2.2 寄存器级计算的载体PE 阵列与寄存器文件的配合所谓寄存器级不是简单说“有个大寄存器堆”而是让每个 PE 除了乘加单元之外还要有足够大的本地寄存器文件并且处理器之间有可配置的互连路径能够让 Q 块广播、K 块流动、PV 结果逐级累加。我用的计算阵列是 32×32共 1024 个 PE。每个 PE 有 64 个 32 位寄存器总计寄存器文件 256KB。这个规模看起来不大但对于 FlashAttention 的分块计算来说已经够用。以 head_dim128 为例一个 64×64 的 S 矩阵 tile 只要 8KB 寄存器空间PV 输出累加器需要 32KB全部塞进 PE 阵列没有任何压力。关键是安排好数据流和存储分配别让本该驻留寄存器的数据中途被挤出到 SRAM。2.3 与 GPU HBM、Chiplet 方案相比有哪些优势方案访存延迟片内带宽能效成本灵活性GPU HBM高跨封装走线中HBM 带宽虽高但延迟大中高高Chiplet 普通封装中中取决于 interposer中高中3D 堆叠 NPU 寄存器级低物理距离短高层间 TSV/混合键合高初期高量产摊薄低这不是说 GPU 方案不好而是任务目标不一样。如果要做通用计算平台GPU 的灵活性无可替代但如果目标是把注意力算子做到极致能效专用 3D 堆叠结构更能把“数据不动、计算动”的设计理念落地。2.4 功耗、散热与良率的现实约束说到 3D 堆叠不能不提发热。逻辑层和存储层摞在一起之后散热面积小了热阻自然上去。业界常用做法是降频但这会抵消掉一部分带宽收益。我在这轮设计里把计算频率定在 1.5GHz 而不是 2GHz 以上保证局部热点不超过 85 摄氏度。同时把参与计算最频繁的 K/V 数据放在靠近散热盖的存储层牺牲一点访问延迟来换可靠性。良率方面大面积的 3D 堆叠会有 TSV 失效的问题所以我在电路里加了冗余数据路径在测试阶段可以屏蔽掉坏点。3. FlashAttention 在寄存器级 NPU 上的映射与实现这一部分是最核心的实操内容。我把 FlashAttention 从算法到硬件一步步拆到块级。3.1 算法变换分块、online softmax 与反向重计算FlashAttention 前向计算有三个关键技巧分块、online softmax、反向重计算。分块的意思是把整个 Q、K、V 矩阵切成 Br×d 和 Bc×d 的小块。online softmax 是指在不知道整行最大值的情况下每处理一个 KV 块就先更新局部最大值 m_i然后按需修正输出 O 的尺度。反向重计算则是在反向传播时不再保存完整的中间注意力矩阵而是用正向时的 keys 和 values 重新算一遍进而节省显存和带宽。硬件实现上有个很微妙的地方GPU 上的 online softmax 通常会把 m_i 和 l_i归一化因子放在寄存器里但每个 block 之间的通信要通过共享内存。NPU 上我让每个 PE 只负责部分行的 softmax 统计量统计量更新通过阵列行方向的归约网络完成整个过程不需要写回 SRAM。实测下来这种编程模型能减少约 30% 的片上同步开销。3.2 块大小的选择给计算阵列配上合适的“胃口”块大小直接决定流水线效率。块取得太大SRAM 放不下就得把中间数据往堆叠 DRAM 搬块取得太小片上带宽利用不满PE 空闲时间长。我按下面的预算来估算head_dim 128输入精度 FP16累加精度 FP32。寄存器中 S 矩阵占用Br × Bc × 2 字节输出累加器占用Br × d × 4 字节核心限制两者之和不超过 256KB 寄存器文件同时 S 块不能跨出 PE 阵列的本地互连范围最终选定 Br64Bc64。这样寄存器中 S 占 8KB输出累加占 32KB还剩大量寄存器给 Q、K、V 的临时数据做缓存。每个 PE 上只要分配 40 字节的固定存储就能容纳这些关键数据几乎没有任何寄存器溢出风险。3.3 寄存器级数据流与伪代码描述下面这段伪代码是我在硬件设计讨论用和实际验证时的基准版本非常接近。它不追求语法上的漂亮重点是把“什么数据留在寄存器里”写清楚。# flash_attention_reg_level: 单头注意力、简化版 # PE 阵列大小 32x32寄存器文件共 256KB # Q, K, V 都已经切成 (num_blocks_q, Br, D) 等形状 m_i [-inf] * Br l_i [0] * Br O_acc zeros(Br, D) # 留在 PE 寄存器累加器中 for bi in range(num_blocks_q): qi load_Q_tile(bi) # 把 Q tile 广播到 PE 阵列 for bj in range(num_blocks_kv): kj load_K_tile(bj) # 从 3D 堆叠 SRAM 中流式读取 vj load_V_tile(bj) s_tile matmul(qi, kj.T) # 分数 tile 直接在寄存器中生成 apply_causal_mask(s_tile) # 因果 mask 也是寄存器级操作 m_prev m_i m_new row_max(m_prev, row_max(s_tile)) alpha exp(m_prev - m_new) p_tile exp(s_tile - broadcast(m_new)) * (s_tile -inf) l_i l_i * alpha row_sum(p_tile) O_acc O_acc * alpha[:, None] matmul(p_tile, vj) m_i m_new # 关键O_acc、m_i、l_i 始终留在 PE 寄存器中不写回 SRAM store_output(O_acc / l_i[:, None])注意代码里 O_acc、m_i、l_i 的量级都留在寄存器里持续更新直到整个序列处理完才写出去。这就是“寄存器级”的关键不是把结果放在 L1而是放在算完就立刻能用的地方。3.4 为什么分块能减少 3D 堆叠层的带宽压力3D 堆叠虽然带宽大但也不是无限大。最好还是尽量复用已经在片上的数据。当 Br64、Bc64 时一个 KV 块只需要从 SRAM 中加载 64×128×2×232KB 数据就能完成 64×64 个分数计算和对应的 PV 累加计算与访存比大约是 128:1。这样即使 3D 堆叠存储层的带宽到不了单颗 HBM 的水平也不用担心成为瓶颈。3.5 实现中要注意的几个硬件设计细节因果 mask 不要用“把分数设成极大负数再 exp”的方式太浪费 PE 计算资源。我用的办法是让分数生成时直接在控制位上跳过非法位置的乘加只有 exp 阶段对这些位置填 0节省约 20% 功耗。矩阵乘的顺序上我选择先算“O_acc * alpha”再做 matmul(p_tile, vj)。这样避免每个 V 块都先乘 alpha减少一次整块数据缩放操作。虽然数学上一样时序上却差不少。PE 间通信要尽量做在行方向上。online softmax 的 row_max 和 row_sum 需要跨列归约如果把归约路径设计成列方向会导致不同头head之间互相干扰流水线很容易卡死。4. 软硬件协同怎么让这块 NPU 真正跑起来硬件设计再好最终还是要接进模型训练和推理框架里。这个环节我发现至少一半的坑不在 RTL 逻辑而在工具链和环境上。4.1 数据搬运引擎硬件算得快但搬数慢照样白搭寄存器级方案省掉了中间结果搬移但你还是要从外部把 Q、K、V 搬进 PE。搬数这件事我用异步 DMA 来完成主计算流水线在做第 bj 个 KV 块的矩阵乘时DMA 同时在预取第 bj1 个 KV 块。DMA 的地址计算逻辑直接内置块状迭代器支持二维 stride 访问这样 K 和 V 可以按自然 layout 流式加载不用在 SRAM 里做重排。没有这个预取机制片上计算阵列的空置率会非常高3D 堆叠的带宽优势根本体现不出来。4.2 集成到 PyTorch 时最常见的环境问题做算法出身的朋友通常把 NPU 当成 CUDA 用结果第一个报错就懵了。最常见的是这条npu is selected as device, but torch_npu is not available. Please ensure torch_npu is installed.这个错误我见过太多次原因基本是三选一没有安装 torch_npu 适配层安装了但版本和当前 PyTorch 不匹配认证环境里DEVICE_ID和驱动版本对不上。排查顺序我建议先从版本匹配查起再看是不是有多个 Python 环境混用。很多时候不是真的没有 torch_npu而是 pip 包装到了另一个环境导入时自然就失败了。4.3 与 Megatron/Swift 等训练框架配合的参考路径在训练框架侧FlashAttention 不会单独存在。它通常作为 attention 模块的内核隐藏在 Megatron 的 context parallel 或者 Swift 的微调流水线后面。我这里走的思路是先做自定义算子把 flash attention 内核封装成torch.autograd.Function然后在 Megatron 的CoreAttention里用环境变量切换到底层实现。Swift 微调场景也是类似的接法针对 LoRA 这类参数高效的场景注意力主路径不变LoRA 部分只作用在 Q 和 V 的输出上算子层面不用改动。4.4 学习这套设计可以看什么教材做一颗专用 NPU光看论文不够。我自己的阅读路径是先精读《计算机体系结构量化研究方法》里关于访存层次和数据流的部分再找 AI 芯片公开课和几本讲 ASIC 设计方法的教材对照着看。如果只是做算法侧优化看 FlashAttention 原始论文加一两篇实现文章就够了但要做到寄存器级和 3D 堆叠这一层必须把体系结构、数字 IC 设计和编译器调度三个领域都打通。5. 性能评估既要算得快也要算得巧这一章讲怎么验证方案到底值不值。性能评估不是只看跑分还要看到底把功耗花在了哪里。5.1 核心指标能效、吞吐、延迟、面积我从五个维度评估这套设计能效TOPS/W注意力算子在特定位宽下的实际有效算力除以功耗。吞吐token/s在给定 batch 和序列长度下的端到端吞吐。延迟ms/token单次生成首 token 的延迟推理场景更关注这个。面积效率PE 阵列和寄存器文件单位面积能产出多少有效算力主要看布局有没有浪费。精度通过对比原始实现和优化实现的 logits 差值确认优化没有破坏数值稳定性。5.2 怎么设计公平对比我做对比时坚持三条原则。第一基线一定也要优化到位不能拿一个普通注意力实现来凑数那样对比没意义。第二误差必须记录FlashAttention 本身因为 online softmax 会有一定的数值差异只要相对误差在 1e-3 以内我认为可以接受。第三比较范围要清晰只比较注意力算子本身还是把模型整条链路也算进去结论完全不同。汇报时必须写清楚。5.3 现实预期提升不会魔法般地出现在仿真环境中这个方案针对长序列seq_len 大于 8192的注意力比“普通 NPU 实现 外部 DRAM 中间结果”的能效要高出 2-3 倍。这个倍数不是算力堆出来的而是把原来写进 SRAM、再读回来做 softmax 的流量砍掉之后的结果。我经常见人只关注 MAC 利用率但实际上对这类访存密集算子能效上最大头的开销是数据搬运本身。寄存器级方案直接把数据搬运压缩到最小单位能效自然就上来了。6. 踩坑记录与实操建议这部分是我最想分享的也是过去几个月被反复折腾出来的经验。6.1 常见问题速查表现象可能原因排查方向仿真跑 10 分钟就卡死DMA 与计算阵列之间的同步没做好检查 DMA 预取回卷标志和依赖记录softmax 数值溢出到 NaN累计变量更新顺序写反复核 m_new 与 l_i 的前后顺序算力利用率不到 50%KV 块太小PE 空转严重增大 Bc减少外循环次数因果 mask 结果错误mask 作用位置不对确认 mask 是在 exp 之前生效torch_npu 导入失败版本不匹配或环境混乱按 4.2 的顺序查版本、路径、设备 ID6.2 三个排障心法第一必要时候先退到纯软件实现。遇到硬件行为不符合预期我先用 Python 写一套完全等价的基准把每一步中间结果打出来再和硬件仿真输出逐块对比很快就能定位是哪段逻辑出了问题。第二先复现再优化。不要一上来就挑战最大序列长度。我把 seq_len 从 256 起步单头单卡跑通再往多头、长序列和训练并行上扩展。每次扩展只改一个变量出问题的时候才不会像无头苍蝇。第三给仿真加“断点”。我在 RTL 仿真里加了几个监控点一旦检测到 S_tile 的非零元素数量偏差超过阈值就立刻停止。这有点像软件开发里的断言能省下大量深夜排查时间。6.3 给后来者的几条具体建议如果让我重新做一遍我会在第一天就搭好“算法参考实现 硬件仿真 parser 对照”三件套而不是先闷头调硬件。硬件设计和算法优化是互相咬合的先把 FlashAttention 的参考实现吃透再动硬件事半功倍。不要一上来就是 RTL 或 PDK 的细节那些等架构确定后再补完全来得及。另外寄存器级优化不是寄存器越多越好。寄存器文件的读写端口和面积都是成本关键是找到“中间结果从生成到消费之间最长的那条路径”把这条路径上的数据留在寄存器里比什么都重要。大多数时候你只需要优化那条最拥挤的数据路径就够了其他地方保持常规 SRAM 访问即可省下的面积和功耗反而可以用来提高并行度。最后分享一个我现在还在用的怪招给每个 KV 块加一个“脏标记”。当一个 KV 块在寄存器中生成后如果发现它在三个计算周期内没有被消费就直接发一个异常宁可让流水线停一拍也不让这种数据在寄存器文件里赖着不走到处占地方。这个设计刚加的时候大家觉得浪费后来发现它能很有效地暴露一些隐藏的数据依赖问题成了调试阶段的“照妖镜”。