前几篇我们把状态空间模型从连续系统一路讲到了Mamba的选择性机制,模型设计层面的故事基本讲完了。但我一直觉得,真正让Mamba在LLM领域站住脚的,不是那个精妙的input-dependent选择想法本身,而是它背后那套工程:并行扫描和硬件感知优化。这篇是这个系列的第八篇,也是Mamba专题的下半部分,我们来把这两块硬骨头彻底啃透。
先说我自己的亲身经历。大概一年多前,我在一台V100上写了一个naive的SSM训练脚本——很简单,直接用for循环把公式 h_t = A_bar * h_{t-1} + B_bar * x_t 一步一步算下去。序列长度设到8192,hidden 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微秒,L=8192时光是启动开销就是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)是线性时不变系统(LTI),A、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, a+b, a+b+c, a+b+c+d]。朴素写法是一个循环,逐个累加,复杂度O(L),这就是"串行扫描"。
但如果把前缀和看成一种特殊的递归,它有一个关键性质:加法满足结合律。因为结合律,你可以先把数组拆成两半,各自并行计算局部前缀和,再把两半连接起来:
- 左半部分直接算:a, a+b
- 右半部分先算:c, c+d
- 然后把左半部分的总和(a+b)加到右半部分每个结果上:c+(a+b), c+d+(a+b)
这样只需要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_2:
h_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-Steele | O(L log L) | O(log L) | 并行度高,实现简单,适合GPU |
| Blelloch | O(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_state=16,d_model=4096,每个时间步的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_state=16,d_model=4096时,一个时间步的state buffer约128KB(fp16)。而一个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), combine=ssm_combine, dim=1, )注意这里我在combine里对A做了unsqueeze(-1)来适配D维的广播。教学版这么写很方便,但性能不是最优。
如果是chunked版本,大致骨架长这样:
def chunked_selective_scan(A_bar, Bx_bar, chunk_size=64): # A_bar: (B, L, N),Bx_bar: (B, L, N, D) B, L, N, D = Bx_bar.shape h = torch.zeros(B, N, D, device=A_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), combine=ssm_combine, dim=1, ) # 叠加上一区块传入的初始状态 # A_prefix = cumprod(A_chunk),表示初始状态经过前i步后的衰减系数 A_prefix = torch.cumprod(A_chunk, dim=1) # (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, dim=1)这个实现虽然性能没法跟官方CUDA比,但逻辑足够清晰,跑小规模实验完全够用。
提示:torch.associative_scan要求combine满足结合律。浮点数运算不严格满足结合律,不同并行顺序会带来极微小的数值差异。训练中使用并行扫描,推理时如果用不同的串行顺序计算,可能出现推理结果与训练时不完全一致的情况。
4.3 坑1:fp16下的数值稳定性和一致性
并行扫描最隐蔽的坑是半精度下的数值问题。Mamba的A_bar = exp(Δ * A),由于A初始化为负值、Δ为正,A_bar会落在(0, 1]区间。当A的元素非常接近0时,A_bar接近1,状态几乎不衰减;长时间递推后,微小误差会被不断放大。
我在fp16下跑过一个d_state=64的实验,发现扫描长度超过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 几个实测记录
下面是我在一张消费级显卡上做的粗略对比,示意数据,环境不同结果会不一样,但趋势可以参考。
| 实现方式 | 序列长度8192,hidden 1024 | 说明 |
|---|---|---|
| PyTorch for循环 | 大约1-2秒/step | 主要耗在kernel launch和HBM反复读写 |
| torch.associative_scan | 大约20-40ms/step | 并行扫描生效,但associate overhead较大 |
| Triton融合kernel | 5-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和矩阵乘法混合使用。它的做法是把序列切成大的block,block内部用矩阵乘法处理(这部分可以走Tensor Core),block之间用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上的,往往比先看它在数学上多精巧更有价值。