
前几篇我们把状态空间模型从连续系统一路讲到了Mamba的选择性机制模型设计层面的故事基本讲完了。但我一直觉得真正让Mamba在LLM领域站住脚的不是那个精妙的input-dependent选择想法本身而是它背后那套工程并行扫描和硬件感知优化。这篇是这个系列的第八篇也是Mamba专题的下半部分我们来把这两块硬骨头彻底啃透。先说我自己的亲身经历。大概一年多前我在一台V100上写了一个naive的SSM训练脚本——很简单直接用for循环把公式 h_t A_bar * h_{t-1} B_bar * x_t 一步一步算下去。序列长度设到8192hidden size只有512我以为这个O(N)的模型应该秒出结果。结果一个forward step跑了上百毫秒反向传播更是让人崩溃。当时我第一反应是这模型是不是白设计了。后来我才意识到问题不在模型而在我完全不理解GPU是怎么执行递归计算的。本文要讲的并行扫描与硬件感知优化正是当年把我从坑里捞出来的两门手艺。1. 为什么SSM的递归计算在GPU上跑不动1.1 GPU的并行胃口和递归的天生矛盾先聊聊GPU这个硬件。GPU本质上是SIMT架构成百上千个线程同时执行相同的指令处理不同的数据。它最喜欢的计算类型是矩阵乘法这种数据并行任务——A100上几千个SM流式多处理器同时算不同的分块互不依赖吞吐能打到非常夸张的程度。但递归计算正好相反。h_t f(h_{t-1})这个链条每一步都依赖上一步的结果物理上无法把第100步和第50步同时算出来。在GPU上写一个长度为L的朴素递归意味着L次顺序的kernel调用每次kernel就算那么一点点矩阵运算。单次kernel launch overhead大约是5-10微秒L8192时光是启动开销就是40-80毫秒还谈不上任何实际计算。这就是我当年那个脚本慢得离谱的根本原因。我见过很多朋友第一次跑SSM时都遇到同样的困惑以为是模型结构有问题其实是执行方式的问题。串行递归和GPU并行架构之间的鸿沟是Mamba这类模型必须解决的第一道工程问题也是并行扫描算法登场的直接动机。1.2 O(N)复杂度在GPU上的陷阱很多人在理解序列模型复杂度时有个误区Transformer是O(N²)Mamba是O(N)所以Mamba应该天生更快。这个说法在算法理论上没有错但在硬件上过于天真。Transformer的O(N²)是矩阵乘法的复杂度而矩阵乘法恰恰是GPU最擅长的事。A100的FP16算力大约是312 TFLOPS即便做一个巨大的注意力矩阵算力利用率也可以非常高。Mamba的O(N)如果退化成L次串行小矩阵运算那算力利用率可能只有个位数百分比——一次kernel里就几百个浮点操作大部分时间花在数据搬运和kernel launch上。这就是I/O bound和compute bound的区别。朴素RNN式执行是典型的I/O bound计算少、搬运多、启动频繁。而优化后的目标是让Mamba在长序列上真正逼近compute bound把算力芯片吃满。1.3 更本质的问题选择性机制毁掉了卷积捷径这里要补充一个系列前面讲过的关键点。经典SSM比如S4是线性时不变系统LTIA、B、C、Δ都是固定参数。LTI系统有一个特别好的性质它对整个序列的作用可以等价成一个全局卷积。既然等价于卷积就可以用FFT或者预计算卷积核的方式实现O(N log N)的并行训练不需要逐时间步递推。但Mamba为了获得选择性——即让模型根据当前token的内容决定遗忘什么、记住什么、输出什么——把B、C、Δ都变成了输入的函数。系统从时不变变成时变卷积等价性不再成立。你没办法预计算一个全局卷积核因为每个token的卷积核都是不一样的。这意味着Mamba必须回到最原始的递推计算。计算复杂度确实是O(N)但执行方式从并行卷积退化成了串行扫描GPU效率暴跌。Mamba论文里最硬核的部分就是把这条看似必然很慢的路用并行扫描和硬件感知优化重新修成了高速路。2. 并行扫描把递归的链条变成并行的大树2.1 前缀和就是并行算法的起点要理解并行扫描最好的切入点是前缀和prefix sum问题。给定一个数组[a, b, c, d]要求输出[a, ab, abc, abcd]。朴素写法是一个循环逐个累加复杂度O(L)这就是串行扫描。但如果把前缀和看成一种特殊的递归它有一个关键性质加法满足结合律。因为结合律你可以先把数组拆成两半各自并行计算局部前缀和再把两半连接起来左半部分直接算a, ab右半部分先算c, cd然后把左半部分的总和(ab)加到右半部分每个结果上c(ab), cd(ab)这样只需要O(log L)轮并行操作而不是L轮串行操作。前缀和的这种并行版本就是scan扫描算法的原型。2.2 SSM递推为什么能被scan现在回到SSM的核心递推式h_t A_t * h_{t-1} B_t * x_t如果只是看这个式子它并不像加法那样显然可并行。但如果我们换个角度把一个时间步的变换看成一个操作符它包含两部分缩放A_t作用于状态和添加项B_t * x_t。定义二元组 o_t (A_t, b_t)其中 b_t B_t * x_t。这个操作符有一个非常好的性质任意两段连续时间步可以合并成一个等价操作符。假设我们有两个相邻的操作符 o_1 (A_1, b_1) 和 o_2 (A_2, b_2)先执行o_1再执行o_2h_1 A_1 * h_0 b_1 h_2 A_2 * h_1 b_2 A_2 * (A_1 * h_0 b_1) b_2 (A_2 * A_1) * h_0 (A_2 * b_1 b_2)所以合并后的操作符是(A_comb, b_comb) (A_2 * A_1, A_2 * b_1 b_2)这个合并操作显然满足结合律先合并o_1、o_2再合并o_3和先合并o_2、o_3再合并o_1最终得到的等价操作符完全一样——因为本质上它们都是在定义从初始状态到末尾状态的线性变换。有了结合律我们就可以像并行前缀和一样把L个时间步的操作符构建成一棵并行合并的大树每轮合并操作数减半L步的串行依赖被压缩成O(log L)轮并行操作。这是并行扫描能够应用在SSM上的数学核心。我当年第一次看懂这个并合规则时觉得非常优雅Mamba的线性递推结构刚好开放出了这一条并行化的路。2.3 两种经典并行扫描Hillis-Steele与Blelloch在GPU上实现scan主要有两个经典算法理解它们有助于明白为什么Mamba实际实现不直接照搬某一种。Hillis-Steele算法思路很直观每一轮每个位置都和前面距离2^k的位置合并一次。经过log L轮后每个位置都拿到了它之前所有元素合并的结果。它的总工作量为O(L log L)但并行度极高很适合GPU这种人多力量大的硬件。Blelloch算法则更精巧分up-sweep和down-sweep两个阶段。up-sweep阶段自底向上构建局部前缀和树down-sweep阶段自顶向下用兄弟节点的和来补齐每个位置的前缀值。它的总工作量是O(L)但常数更大并行度比Hillis-Steele低。算法总工作量并行步数特点Hillis-SteeleO(L log L)O(log L)并行度高实现简单适合GPUBlellochO(L)O(log L)工作量最优但常数大、同步多实际工程里Mamba的selective scan并没有纯用某一个算法而是先分块块内用scan、块间用串行递推并把扫描和后续输出计算融合进同一个kernel。2.4 Mamba实际怎么切块chunked scan理论上并行扫描可以覆盖整个序列长度但实际GPU kernel不能对无限长的序列直接暴力扫描。原因是线程块thread block能用的SRAM片上内存是有限的我们需要在块内保存各个时间步的中间状态。处理超长序列时更现实的做法是把序列切成若干chunk每个chunk内部用并行扫描chunk之间保持串行依赖。这其实是一个很自然的混合策略chunk内部的扫描是并行的但每个chunk需要等上一个chunk的最终状态传进来才能开始计算下一个chunk。序列越长chunk数量越多串行部分占的比例越大。好消息是chunk内部的并行层数已经是log(chunk_size)级别即使是8192的序列切分成64长度的chunk整体串行步数也只有128步左右远好于原始8192步。chunk size的选择不是随便定的它受限于SRAM容量和状态维度的大小。假设d_state16d_model4096每个时间步的state buffer是16×4096×2字节fp16约128KB。而A100每个SM的SRAM只有192KB左右。这就意味着一个chunk里最多也就缓存1-2个时间步的完整状态chunk不可能设得很大。这部分细节我们在下一章硬件感知优化里再展开。3. 硬件感知优化把时间省在内存层级上3.1 GPU内存层级和带宽差的真相很多写PyTorch的人对内存层级不太敏感以为GPU只有显存这一种内存。实际上GPU内部至少有两层内存全局显存HBM高带宽内存和片上SRAM也叫共享内存。SRAM容量小A100单SM约192KB但带宽极高HBM容量大几十GB但带宽远低于SRAM。具体数字会因GPU型号不同有差异大致量级是HBM带宽约2TB/s而SRAM的聚合带宽可以达到十几TB/s甚至更高。更重要的是数据从HBM搬到SRAM或者从SRAM搬回HBM都需要显式的load/store指令并且有延迟。深度学习计算的本质是从HBM取数据到SRAM/寄存器计算完再写回去。如果一次计算需要反复读写HBM那么再快的算力也被内存带宽卡死。这就是为什么有些kernel慢不是ALU不够快而是数据搬运占了绝大部分时间。一个非常直观的类比假设你是厨师HBM是楼下的大仓库SRAM是厨房台面。如果每炒一个菜都要跑一趟楼下仓库取食材那你做菜再熟练也没用时间全花在路上了。优化目标就是一次下楼把够做一整桌菜的食材都搬上来在台面上完成所有处理。3.2 Kernel Fusion把多次kernel调用合成一次Mamba的核心递推如果要朴素实现每一步时间步至少涉及离散化计算A_bar和B_bar、状态更新h_t A_bar * h_{t-1} B_bar * x_t、输出投影y_t C_t * h_t三个操作。每一步都会在HBM上读写中间结果而每个中间结果的写入都需要带宽。硬件感知优化的第一件事就是kernel fusion把这些操作合并到一个kernel里让中间结果留在SRAM不落回HBM。selective scan的整个循环以及所有的逐元素操作都被融合进一个CUDA kernel中。这样前向传播过程中输入和输出各读一次、写一次HBM中间状态全部留在片上。这套思想和FlashAttention一脉相承。FlashAttention也是把attention的计算融合成单次kernel避免把巨大的注意力矩阵写回HBM。Mamba论文明确说了它从FlashAttention的工程思路里借鉴了很多。3.3 Selective Scan的重计算策略用算力换内存并行扫描解决了计算效率问题但反向传播还有一个大坑求梯度需要各个时间步的中间状态h_t。如果每个时间步都保存一份state buffer内存占用是O(L × d_state × d_model)在长序列和宽模型下直接爆炸。Mamba的选择是不保存中间状态反向传播时重新计算。这听起来像PyTorch的gradient checkpointing但Mamba是在kernel内部做的recompute——forward时只保存每个chunk边界的状态反向时利用边界状态和保存的输入在短时间内重新跑一遍扫描再计算梯度。这就是典型的用算力换内存。如果扫描本身够快尤其并行扫描让重算代价变得很低这个交换就非常划算。我实测下来Mamba训练内存可以压到接近单层Transformer的水平这条recompute策略功不可没。3.4 一个具体的尺度估算chunk size为什么不能大回到上一章的悬念。selective scan在GPU上的典型状态空间是d_state×d_model。d_state16d_model4096时一个时间步的state buffer约128KBfp16。而一个SM的SRAM上限通常只有100-200KB。如果你还希望chunk里缓存多个时间步的中间值chunk size必须小。这就是为什么Mamba的selective scan通常按chunk处理并且chunk的大小不是看序列长度而是被state维度卡住。换句话说Mamba在GPU上的效率瓶颈是state buffer的片上缓存能力而不是序列长度。这也解释了为什么Mamba 2要把state维度设计得更加紧凑并用chunked矩阵乘法进一步压榨Tensor Core——这部分我们放到第五章展开。4. 把Mamba跑起来复现路径、工程细节与踩坑记录4.1 代码选择官方CUDA实现 vs 教学PyTorch实现如果你想自己跑Mamba首先面临代码选型。官方仓库state-spaces/mamba提供的是高性能CUDA实现速度没问题但如果你想读懂它并做定制那几千行CUDA代码属实劝退。教学向的话推荐mamba.py一个单文件的PyTorch实现可读性极强以及各种Triton版本的实现。Triton的好处是用Python写GPU kernel比CUDA容易上手得多而且性能已经相当不错。我的建议是分两步走先用教学版把机制跑通理解scan和selective的完整逻辑再切到官方实现训练大模型。直接上官方代码做研究调试成本会很高因为你很难判断一个bug是模型逻辑错了还是kernel写崩了。4.2 一个正确的associative scan combine实现如果你要在PyTorch里自己实现并行扫描核心就是那个combine算子。这里给一个可读性优先的示意版本注意状态维度之间的乘法是逐元素操作对应Mamba的对角化状态设计import torch def ssm_combine(op1, op2): # 每个op是一个二元组 (A, bu)表示 h A * h bu # A: (..., N)bu: (..., N, D) A1, bu1 op1 A2, bu2 op2 # 合并先执行op1再执行op2 A_comb A2 * A1 bu_comb A2.unsqueeze(-1) * bu1 bu2 return A_comb, bu_comb配合torch.associative_scan就可以对整段序列做并行扫描# A_bar: (B, L, N)Bu_bar: (B, L, N, D) h_local torch.associative_scan( (A_bar, Bu_bar), combinessm_combine, dim1, )注意这里我在combine里对A做了unsqueeze(-1)来适配D维的广播。教学版这么写很方便但性能不是最优。如果是chunked版本大致骨架长这样def chunked_selective_scan(A_bar, Bx_bar, chunk_size64): # A_bar: (B, L, N)Bx_bar: (B, L, N, D) B, L, N, D Bx_bar.shape h torch.zeros(B, N, D, deviceA_bar.device) outputs [] for start in range(0, L, chunk_size): end min(start chunk_size, L) A_chunk A_bar[:, start:end] # (B, C, N) Bx_chunk Bx_bar[:, start:end] # (B, C, N, D) # 块内并行扫描零初始状态 local_h torch.associative_scan( (A_chunk, Bx_chunk), combinessm_combine, dim1, ) # 叠加上一区块传入的初始状态 # A_prefix cumprod(A_chunk)表示初始状态经过前i步后的衰减系数 A_prefix torch.cumprod(A_chunk, dim1) # (B, C, N) full_h local_h A_prefix.unsqueeze(-1) * h.unsqueeze(1) h full_h[:, -1] outputs.append(full_h) return torch.cat(outputs, dim1)这个实现虽然性能没法跟官方CUDA比但逻辑足够清晰跑小规模实验完全够用。提示torch.associative_scan要求combine满足结合律。浮点数运算不严格满足结合律不同并行顺序会带来极微小的数值差异。训练中使用并行扫描推理时如果用不同的串行顺序计算可能出现推理结果与训练时不完全一致的情况。4.3 坑1fp16下的数值稳定性和一致性并行扫描最隐蔽的坑是半精度下的数值问题。Mamba的A_bar exp(Δ * A)由于A初始化为负值、Δ为正A_bar会落在(0, 1]区间。当A的元素非常接近0时A_bar接近1状态几乎不衰减长时间递推后微小误差会被不断放大。我在fp16下跑过一个d_state64的实验发现扫描长度超过4096后某些state维度的值开始出现明显漂移最终导致loss不降。排查后发现不是模型设计问题而是半精度的累积误差。解决办法是在A的初始化上动手脚保证初始A在负方向有一定幅度不要初始化为0附近同时必要时对中间状态做逐chunk的clamp。另一个一致性坑是训练时你用了并行scan推理时如果图省事写了串行循环两者虽然数学等价但浮点运算顺序不同结果会有一点点差异。对于自回归生成这种差异可能会随步骤累积出现微妙的行为偏移。我的习惯是训练和推理共用同一套scan kernel避免这类不可复现的玄学bug。4.4 坑2朴素的O(N)实现没有意义我把话放在这里如果只是把Mamba的循环丢进PyTorch里跑它的速度很可能比同规模的Transformer还要慢。原因我们第一章已经分析过——串行循环、kernel launch开销、中间张量反复落HBM。所以做性能对比实验时不要拿朴素版本Mamba和优化过的Transformer对比然后得出Mamba不行的结论。要测Mamba的真实性能至少使用Triton实现或官方CUDA kernel。在我自己的对比里同一模型规模官方kernel的吞吐大约是朴素PyTorch实现的几十倍。硬件感知优化不是锦上添花而是Mamba能够成为LLM架构候选的前提条件。4.5 几个实测记录下面是我在一张消费级显卡上做的粗略对比示意数据环境不同结果会不一样但趋势可以参考。实现方式序列长度8192hidden 1024说明PyTorch for循环大约1-2秒/step主要耗在kernel launch和HBM反复读写torch.associative_scan大约20-40ms/step并行扫描生效但associate overhead较大Triton融合kernel5-15ms/step更接近可用的生产性能5. 从Mamba 1到Mamba 2硬件感知优化的下一步5.1 Mamba 2的发现状态空间和注意力之间的对偶Mamba 2论文标题叫《Transformers are SSMs》中文直译Transformer就是状态空间模型。这个标题背后是一个漂亮的数学发现SSM的线性递推可以被展开成一个半分离矩阵semiseparable matrix的矩阵乘法而这个矩阵乘法和某些形式的线性注意力有着对偶关系。这个发现的工程意义非常直接既然SSM能被表达成矩阵乘法那就可以用GPU上高度优化的矩阵乘法库cuBLAS、Tensor Core来加速而不只是依赖手写的scan kernel。Mamba 1的selective scan虽然已经很快但它的瓶颈在于scan本质上是记忆体密集的逐元素操作无法充分利用Tensor Core。Mamba 2通过block矩阵分解把大部分计算量转成了大矩阵乘法GPU利用率大幅提升。5.2 chunked scan在Mamba 2里的进一步演进Mamba 2并没有完全抛弃scan而是把scan和矩阵乘法混合使用。它的做法是把序列切成大的blockblock内部用矩阵乘法处理这部分可以走Tensor Coreblock之间用scan处理。这和我们在Mamba 1里讲的chunked scan思路一脉相承只是把chunk内部的并行单元从逐元素scan换成了矩阵乘法粒度更粗、效率更高。带来的直接好处是Mamba 2可以使用更大的d_state比如128甚至256而不像Mamba 1那样被state buffer的SRAM容量死死卡住。更大的状态维度意味着更强的记忆能力这也是Mamba 2在部分基准上超过Mamba 1和部分Transformer变体的原因之一。如果你把Mamba 1的硬件感知优化理解成在SRAM里精打细算地做扫描那Mamba 2的思路就是尽量让更多计算变回矩阵乘法只在不得不用scan的地方保留scan。这个迁移思路值得所有做模型优化的人学习。5.3 硬件感知工程思想正在成为序列建模的共识观察Mamba、FlashAttention、RWKV、线性注意力等一系列工作会发现一个清晰的趋势单纯设计数学上更高效的模型已经不够了必须让模型的计算模式匹配硬件的存储层次和执行模型。三条通用原则可以说是这几年的经验总结第一优先把计算组织成矩阵乘法这是GPU最擅长的事。第二凡是涉及序列维度的操作优先考虑能否用结合律做并行归约scan避免逐时间步串行。第三kernel融合和recompute的组合拳几乎适用于所有长序列算子不要在Python层反复生成中间张量。我自己在实现一个新算子前现在会先问三个问题这个计算能不能分块块的边界在哪中间状态能不能不落HBM大部分性能问题这三个问题想清楚了都能解决一大半。如果你只读这个系列的最后两篇我的建议是重点理解并行扫描背后的结合律思维和硬件感知里的内存搬运算力视角。Mamba的具体结构会过时工程思想不会。下次遇到任何一个新的序列模型先看它是怎么跑在GPU上的往往比先看它在数学上多精巧更有价值。