Mamba 选择性状态空间模型:亚二次复杂度大语言模型推理工程实践

从 S4 到 Mamba 的架构演进,选择性状态空间如何实现 Transformer 级别的建模能力与 RNN 级别的推理效率?本文深入剖析 SSM 的数学本质、选择性机制的硬件感知实现,以及 Mamba 在百亿参数规模的工程部署实践。

一、Transformer 推理的根本瓶颈

自 2017 年 Attention is All You Need 发表以来,Transformer 凭借其卓越的并行训练能力和全局依赖建模,成为大语言模型的事实标准。然而,Transformer 的自注意力机制存在两个核心效率瓶颈:

训练阶段:注意力矩阵的计算复杂度为 O(N²),其中 N 为序列长度。虽然 Flash Attention 在 GPU HBM 与 SRAM 之间做了精巧的 IO 感知分块,但计算量本身并未减少。

推理阶段:自回归生成时,模型需要维护一个 Key-Value Cache(KV Cache),其内存占用与序列长度呈线性增长。对于 70B 模型、8K 上下文长度,KV Cache 可达数 GB,严重制约了推理吞吐和批处理能力。

业界已经提出了诸多方案试图缓解这一问题:线性注意力(Linear Attention)牺牲了精确记忆能力;稀疏注意力(如 BigBird、Longformer)仅适用于特定任务;RWKV 尝试用线性时间递推的 RNN 结构替代注意力。而 2023 年 12 月提出的 Mamba(Gu & Dao, 2023)则走出了一条全新路线——选择性状态空间模型(Selective State Space Model, SSM),在保持 Transformer 级建模能力的同时,实现了推理阶段的线性复杂度。


二、状态空间模型:从控制论到深度学习

2.1 SSM 的连续时间形式

状态空间模型起源于控制论,描述一个连续时间动态系统:


h'(t) = A · h(t) + B · x(t)
y(t) = C · h(t) + D · x(t)

其中 x(t) 是输入信号,h(t) 是隐藏状态(latent state),y(t) 是输出,矩阵 A 控制状态演化,B 控制输入如何注入状态,C 控制状态如何映射到输出,D 是直连跳跃项(Connection)。

这个 ODE 的含义很直观:h(t) 持续记忆历史输入信息,其演化方式由 A 的特征值决定——适当的 A 设计能使状态在长期依赖中既不死机(梯度消失)也不爆炸(梯度发散)。

2.2 离散化:从连续到神经网络层

在实际深度学习中,我们处理的是离散序列 x₀, x₁, ..., xₜ。通过零阶保持(Zero-Order Hold, ZOH)离散化,步长为 Δ(采样间隔),SSM 变为:


Ā = exp(ΔA)
B̄ = (ΔA)⁻¹ · (exp(ΔA) - I) · B
h_t = Ā · h_{t-1} + B̄ · x_t
y_t = C · h_t + D · x_t

这里有一个关键观察:离散化后的 SSM 在形式上就是一个 RNN。输入 x_t 被累加到隐藏状态 h_t,然后 h_t 作为下一时间步的记忆。与标准 RNN 不同的是,SSM 的 A 矩阵来源于连续时间系统的解析离散化,而非直接学习,这使得状态演化具有更好的数学稳定性和表达能力。

2.3 并行计算的卷积视角

将上述递推展开,可以得到 y_t 关于所有历史输入的显式表达:


y_t = Σ(k=0..t) C · Āᵏ · B̄ · x_{t-k} + D · x_t

这正是一个卷积运算,滤波器核为:


K = (C·B̄, C·Ā·B̄, C·Ā²·B̄, ..., C·Ā^{N-1}·B̄)

因此,SSM 的并行训练可以通过高效卷积实现,使用 FFT 将复杂度从 O(N²) 降至 O(N log N)。这就是 S4(Structured State Space Sequence Model)高效训练的数学基础。


三、Mamba 的核心创新:选择性机制

3.1 S4 的局限性

S4(Gu et al., 2021)通过以下设计取得了突破:

  • A 矩阵采用 HiPPO 初始化的高阶正交结构,最大化长期记忆保持能力
  • NPLR(Normal Plus Low-Rank)参数化确保矩阵幂运算的数值稳定性
  • 对角化结构高效计算矩阵幂

但 S4 有一个致命缺点:模型参数 A, B, C 在所有时间步共享,与输入无关。这意味着 SSM 无法对不同的输入 token 采取不同的策略——它无法判断"这个 token 很重要,我需要更努力地记住它"或"这个 token 无关紧要,我可以快速遗忘它"。

类比注意力机制:Transformer 通过 softmax 计算权重,使得重要 token 获得更多关注。而 S4/SSM 对所有 token 一视同仁地应用固定变换,无法动态地选择记忆内容。

3.2 选择性状态空间

Mamba 的核心创新是让 B, C, Δ 参数变成输入依赖的动态参数:


B_t = Linear_B(x_t)
C_t = Linear_C(x_t)
Δ_t = softplus(Linear_Δ(x_t) + broadcast_Δ)
A = 全局学习参数(不依赖于输入)

Δ 的作用尤其关键——它控制状态更新的"速度":

  • 当 Δ 很大时,Ā = exp(-exp(Δ) · A) 快速衰减,历史记忆被迅速遗忘,模型只关注最近输入
  • 当 Δ 很小时,Ā 接近恒等矩阵,状态缓慢衰减,历史记忆被长期保留

这种选择性赋予 Mamba 两个能力:

  1. 有选择地记住或忽略输入:对于"虽然"、"但是"等转折词后的关键信息,增大 Δ 让模型更快更新状态;对于标点符号、填充词等无关信息,减小 Δ 让状态不发生变化
  2. 信息滤波:在推理时,无关信息被直接从记忆中过滤掉,不会污染上下文
  3. 3.3 选择性机制的直觉解释

    以一个具体例子说明选择性如何工作:

    
    输入序列:"赵明是公司的技术总监,但他昨天提交了辞职报告。"
    

    在处理到"但"时,模型需要抛弃先前的部分记忆来重点关注转折后的内容。Mamba 通过增大 Δt 快速衰减旧状态来实现这一点。而 S4 的固定参数无法感知"这里发生了转折",会无差别地将先前的状态混入后续计算。

    另一个场景是精确记忆:在处理需要长距离回忆的任务时,模型需要精确记住某个实体名称。Mamba 可以通过极小的 Δ 让这个信息几乎不受干扰地保存在状态中——类似 Attention 中的硬 select,但实现方式完全不同。


    四、硬件感知的并行算法

    4.1 并行扫描(Parallel Scan)

    选择性机制引入输入依赖的 B_t, C_t, Δ_t 后,并行卷积方案不再适用。Mamba 需要回归递推形式:

    
    h_t = Ā_t · h_{t-1} + B̄_t · x_t
    

    但这个递推可以通过并行前缀和(Parallel Scan / Blelloch 算法)高效计算。对于结合算子 ⊗ 的定义:

    
    (a₁, b₁) ⊗ (a₂, b₂) = (a₂ · a₁, a₂ · b₁ + b₂)
    

    递推的每一步可以表示为 (h_t, x_t) → (Ā_t · h_{t-1} + B̄_t · x_t, x_t),而并行扫描用分治法在 O(log N) 步内完成整个序列的计算。在 GPU 上这被优化为高效的并行归约操作。

    4.2 内存层级优化

    Mamba 最重要的工程贡献:状态驻留在 SRAM,永不写入 HBM。

    在 Transformer 推理中,KV Cache 必须常驻 HBM(几十 GB 显存),每次前向计算都需要从 HBM 读入 KV。对于长序列,HBM 带宽成为吞吐瓶颈。

    Mamba 的策略:

    • 序列的每个位置对应一个时间步的递推
    • 每个 token 对应一个隐藏状态 h_t,维度较小(典型值 d_state=128,远小于 KV Cache 的维度)
    • 在 kernel 融合中,系统矩阵 (A, B_t, C_t, Δ_t) 和隐藏状态 h_t 全部保留在 GPU SRAM
    • 只在序列处理完成后,才将最终状态 h_N 写回 HBM(用于下一层的输入或下一次推理的初始状态)

    这种设计与 Flash Attention 的思路一脉相承——都是将中间计算结果保留在 SRAM 中,最小化 HBM 读写。但 Mamba 的优势更加彻底:跨 token 的隐藏状态更新完全不需要任何 HBM 访问,而 Flash Attention 仍需读写分块的注意力矩阵。

    4.3 内核融合实践

    Mamba 通过 Triton GPU 编程框架实现了高度 fuse 的内核:

    
    # Mamba 的简化伪代码(选择性扫描内核)
    import triton
    import triton.language as tl
    
    @triton.jit
    def selective_scan_kernel(
        A_ptr, B_ptr, C_ptr, delta_ptr,
        hidden_ptr, x_ptr, y_ptr,
        seq_len, state_dim, block_size: tl.constexpr
    ):
        pid = tl.program_id(0)
        
        # 从 HBM 加载当前时间步的输入和动态参数
        x = tl.load(x_ptr + pid * state_dim + tl.arange(0, state_dim))
        b = tl.load(B_ptr + pid * state_dim + tl.arange(0, state_dim))
        c = tl.load(C_ptr + pid * state_dim + tl.arange(0, state_dim))
        delta = tl.load(delta_ptr + pid * state_dim + tl.arange(0, state_dim))
        
        # 离散化:计算当前步的 Ā 和 B̄
        a_base = tl.load(A_ptr + tl.arange(0, state_dim))  # 全局 A 矩阵
        dA = tl.exp(delta * a_base)
        dB = delta * b
        
        # 递推更新(保持在 SRAM)
        h = tl.load(hidden_ptr + tl.arange(0, state_dim))
        h = dA * h + dB * x
        y = c * h
        
        tl.store(y_ptr + pid * state_dim + tl.arange(0, state_dim), y)
    

    Mamba v2(2024 年 4 月发布)进一步优化,将状态维度增大、并行扫描内核的并行度提升,在 NVIDIA A100/H100 上实现了比 Flash Attention 2 更优的吞吐。


    五、Mamba 在大规模语言模型中的表现

    5.1 语言建模基准

    Mamba-2.8B(2.8B 参数)在 Pile 语言建模基准上,其困惑度(perplexity):

    • 显著优于同参数量的 Transformer 模型(如 Pythia-2.8B、OPT-2.7B)
    • 在长上下文任务中(PG19、Books3),SSM 的困惑度随上下文长度增加而持续下降,而 Transformer 的困惑度在较长上下文后不再改善
    • 大幅超越同规模的 RWKV-3B(约 1.5-2.0 ppl 差距)

    Mamba-1.4B 在零样本推理任务(HellaSwag、PIQA、Winogrande 等)中普遍超越 Pythia-1.4B 和 RWKV-1.5B,在部分任务上(如 ARC-Challenge、Lambada)甚至接近 Transformer-6.9B。

    5.2 推理效率

    Mamba 在推理速度上与 Transformer 的对比(基于官方数据,NVIDIA A100 80GB):

    指标 Transformer-1.3B Mamba-1.4B
    训练吞吐量(tokens/s) ~360K ~380K
    推理延迟(8K上下文) ~45ms ~12ms
    推理显存(8K上下文) ~16GB ~6GB
    生成吞吐(tokens/s) ~22 ~83

    Mamba 的推理延迟与序列长度无关(固定状态维度),而 Transformer 每个生成都需要全量 Attention 计算。这使得 Mamba 在长上下文场景中优势极为突出——8K 上下文时延迟差 3-4 倍,64K 上下文时差距超过 10 倍。

    5.3 Mamba 2 的架构演进

    原始的 Mamba 仍采用 Transformer 风格的 MLP+Mamba 混合架构。Mamba 2(2024)进一步修改了 SSM 结构:

    • 引入了类似 Attention 的双路径结构,增强信息交互
    • 使用更大的状态维度(128→256),提高记忆容量
    • 改进了并行扫描的原生并行度

    Mamba 2-2.7B 在几乎所有语言建模基准上超越了 Llama-3-8B(8B 参数),参数量比仅为 1:3,这是 Mamba 系列最惊人的数据之一。


    六、工程实践:部署与优化

    6.1 快速体验

    目前最成熟的 Mamba 推理实现是 mamba-ssm 库(基于 Triton)和 transformers 库。

    
    import torch
    from transformers import Mamba2ForCausalLM, AutoTokenizer
    
    # 加载模型与分词器
    model = Mamba2ForCausalLM.from_pretrained(
        "state-spaces/mamba2-2.7b",
        dtype=torch.bfloat16,
        device="cuda"
    )
    tokenizer = AutoTokenizer.from_pretrained("state-spaces/mamba2-2.7b")
    
    # 自回归生成
    def generate(prompt, max_new_tokens=256, temperature=0.7):
        input_ids = tokenizer.encode(prompt, return_tensors="pt").cuda()
        
        for _ in range(max_new_tokens):
            with torch.no_grad():
                output = model(input_ids, cache_params=None, stateful=True)
            logits = output.logits[:, -1, :]
            if temperature == 0:
                next_token = logits.argmax(dim=-1, keepdim=True)
            else:
                probs = torch.softmax(logits / temperature, dim=-1)
                next_token = torch.multinomial(probs, num_samples=1)
            
            input_ids = torch.cat([input_ids, next_token], dim=1)
            if next_token.item() == tokenizer.eos_token_id:
                break
        
        return tokenizer.decode(input_ids[0], skip_special_tokens=True)
    
    # 调用
    response = generate("请用简单的语言解释状态空间模型的工作原理。")
    print(response)
    

    6.2 状态缓存与推理服务

    在实际部署场景中,Mamba 的独特状态管理给服务框架带来了新的挑战和机会。

    状态压缩与池化(State Pooling)

    Mamba 的状态是一个维度固定的向量(如 d_state=256),不需要按 token 分页。可以采用:

    • 前缀状态缓存(Prefix State Caching):在处理相同前缀的多个请求时直接复用已计算好的层状态
    • 连续批处理优化:对于新到达的 token,只需增量递推一步,无需从头计算
    
    # 概念性的状态缓存复用示例
    class MambaStateCache:
        """缓存不同前缀对应的 Mamba 层状态"""
        def __init__(self, model_config, max_cache_size=1024):
            self.cache = {}  # prefix_hash -> [layer_states]
            self.max_size = max_cache_size
        
        def get(self, input_ids: torch.Tensor):
            """获取已缓存的状态"""
            key = hash(input_ids.tolist())
            if key in self.cache:
                return self.cache[key]
            return None
        
        def put(self, input_ids: torch.Tensor, states):
            if len(self.cache) >= self.max_size:
                # LRU 淘汰
                self.cache.pop(next(iter(self.cache)))
            key = hash(input_ids.tolist())
            self.cache[key] = states
    

    6.3 批处理的加速挑战

    Mamba 推理的一大挑战是动态形状批处理。不同请求的长度和已生成位置不同,导致递推步数不一致。

    目前工程上的解决方案:

    1. 按需前向计算(Continuous Batching):仅在新生 token 上做前向,但同一批内不同请求需独立执行扫描
    2. 定制的 CUDA Graph:针对不同序列长度预编译多个 Graph,按实际长度选择执行
    3. Warp 级并行化:将批内请求的状态排列在连续线程中,利用 GPU SIMT 并行执行
    4. 6.4 与 MoE 架构的互补

      Mamba 与 Mixture-of-Experts(MoE)形成互补:

      • Mamba 解决序列长度维度的效率瓶颈(注意力复杂度)
      • MoE 解决参数总量维度的效率瓶颈(各 token 仅激活部分专家)

      MambaMoE(2024 年提出)将 Mamba 层替换 Attention,保留 MoE 前馈网络。在部分基准上,Mamba-MoE-3B(激活 1.3B)表现接近 Transformer-MoE-8B(激活 4B),展示了极具效率潜力的新架构方向。


      七、Mamba 的局限与当前挑战

      7.1 精确召回能力

      Mamba 通过固定维度的连续状态向量记忆信息,其记忆容量存在上限。对于需要精确区分某个信息是否"在上下文中出现过"的任务(如 key-value 精确检索),Mamba 弱于 Transformer。

      解决方案方向:

      • 混合架构:Jamba(AI21 Labs)、Zamba 等模型在部分层使用全局注意力,其他层使用 Mamba
      • 外部记忆:将精确信息存入外部数据库,由 SSM 控制检索时机
      • 状态维度扩展:增大 d_state(代价是状态不再适合完整驻留 SRAM)

      7.2 超长序列续推

      Mamba 在推理开始时通常从全零状态初始化。扩展到 100K+ 上下文时,需要状态续推(State Continuation)机制——在分块处理超长序列时保持跨块状态的连续性,目前工程实践中仍缺乏标准化方案。

      7.3 硬件生态

      Mamba 依赖定制的 Triton 内核,目前 ONLY 对 NVIDIA GPU(Hopper/Ada/Ampere 架构)有最优支持。AMD GPU(ROCm)和 Apple Silicon(MPS)的支持仍在完善。在 Intel Gaudi、AWS Trainium/Inferentia 等平台上,Mamba 的性能优势不明显。


      八、未来方向与选型建议

      8.1 SSM 与 Transformer 的趋同

      2024 年下半年,SSM 与 Transformer 呈现出有趣的趋同趋势:

      • Attention 架构借鉴 RNN:RetNet、RWKV-6 引入"循环等价"推理模式
      • SSM 架构借鉴 Attention:Mamba 2 的双路径设计本质上类似"内容寻址的注意力"
      • 混合架构涌现:Zamba、Jamba、Griffin 等模型混合注意力与循环结构
      • 硬件层面趋同:各架构都在向"状态驻留 SRAM、最小化 HBM 读写"方向优化

      未来一年,混合架构可能取代纯 Transformer 成为长上下文部署的新标准。

      8.2 选型框架

      场景 推荐架构 原因
      超长上下文推理(32K+) Mamba / 混合架构 延迟与上下文长度无关,显存恒定
      需要精确信息检索 Transformer / 混合架构 状态压缩无法保证精确召回
      边缘设备推理(低显存) Mamba 无 KV Cache 显存压力
      训练新基础模型 Transformer 生态成熟、训练效率稳定
      长文档摘要/代码生成 Mamba 依赖全局模式而非精确 token 召回
      实时流式生成(聊天) Mamba 恒定低延迟

      8.3 基础设施准备

      对于已有 Transformer 部署的团队,引入 Mamba 需要考量:

      1. 模型转换:现有 Attention 模型无法直接转为 Mamba,需从头训练或蒸馏
      2. 推理框架:vLLM 的 Mamba 支持仍在早期阶段,自研内核需一定工程投入
      3. 监控指标:需要新增"状态维度利用率"、"SRAM 与 HBM 读写比例"等性能指标

      4. 结语

        Mamba 不是 Transformer 的完美替代者,而是一个全新的模型架构范式——选择性状态空间。它在序列建模中引入了"内容感知"的稀疏记忆能力,同时在计算层面实现了硬件感知的并行扫描和状态压缩。

        从 S4(结构化状态空间)到 Mamba(选择性状态空间)再到混合架构(如 Jamba、Zamba),模型设计从"结构归一化"向"能力组合化"演进的趋势日益明显。在推理成本和长上下文需求成为核心瓶颈的 2025 年,SSM 家族在工程部署中的价值将愈发凸显。

        对于工程师而言,理解 Mamba 的选择性原理和硬件感知实现,不仅有助于在合适的场景中做出更好的架构选择,也提供了看待"记忆与计算"这一基础命题的新视角——连续时间动力学系统如何高效地离散化为神经网络的计算图,可能是接下来两年深度学习基础设施进化的核心方向之一。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部