2023 年,Albert Gu 和 Tri Dao 发表的论文《Mamba: Linear-Time Sequence Modeling with Selective State Spaces》在深度学习架构领域投下了一颗重磅炸弹。当整个行业还在为 Transformer 的二次复杂度焦头烂额时,Mamba 用 O(n) 的推理复杂度和接近 Transformer 的性能,证明了"替代方案"并非空中楼阁。

一、Transformer 的阿喀琉斯之踵

Transformer 的核心是 Self-Attention 机制。给定长度为 n 的序列,Attention 需要计算所有 token 对之间的相似度,时间复杂度为 O(n²)。这意味着:

  • 序列长度翻倍,计算量翻四倍
  • 推理时 KV Cache 的内存随序列线性增长
  • 长文档处理成本急剧膨胀

虽然 FlashAttention、Ring Attention、MQA/GQA 等技术在工程层面做了大量优化,但 O(n²) 的复杂度本质没有改变。当我们尝试处理 100K、1M 甚至更长序列时,纯粹的工程技巧已触及天花板。

问题的根源不在工程,而在数学。


二、状态空间模型的数学直觉

状态空间模型(State Space Model, SSM)是控制论中的经典工具。其核心思想可以用一个简洁的微分方程描述:

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

其中 x(t) 是输入信号,h(t) 是隐藏状态,y(t) 是输出。离散化后得到递推形式:

h_k = Ā · h_{k-1} + B̄ · x_k
y_k = C · h_k

这个递推结构天然具备两个关键优势:

  1. 线性时间复杂度:每个时间步只需一次矩阵-向量乘法,总体 O(n)
  2. 有界内存:隐状态 h 的维度固定,不随序列长度增长
  3. 但传统 SSM(如 S4)有一个致命缺陷:模型参数 A、B、C 是固定的,对所有输入一视同仁。这意味着它无法像 Attention 那样选择性关注某些 token。


    三、Mamba 的核心突破:选择性状态空间

    Mamba 的第一个关键创新是选择性(Selectivity)。传统 S4 的 Δ、B、C 是常数,Mamba 让它们依赖输入:

    Δ_k = Δ(x_k),  B_k = B(x_k),  C_k = C(x_k)

    这个看似简单的改动带来了质的飞跃:

    # S4: 固定参数,无差别处理
    h_k = A_bar * h_{k-1} + B_bar * x_k  # 对所有 token 一样
    
    # Mamba: 输入依赖,选择性遗忘与记忆
    delta = softplus(W_delta @ x_k)       # 依赖输入的离散化步长
    B = W_B @ x_k                         # 依赖输入的投影
    C = W_C @ x_k                         # 依赖输入的输出投影
    h_k = (1 - delta) * h_{k-1} + delta * (A @ h_{k-1} + B * x_k)

    当某个 token 不重要时,模型可以让 Δ→0,此时 h_k ≈ h_{k-1},相当于跳过了该 token;当 token 重要时,Δ 变大,模型更新状态。这种机制在功能上等价于 Attention 的"选择性关注",但复杂度是线性的。

    3.1 硬件感知的并行算法

    仅靠选择性不够。递推结构虽然理论上线性,但不像矩阵乘法那样天然适合 GPU 并行。Mamba 作者设计了一个硬件感知(Hardware-Aware)的并行算法:

    # 将递推转化为并行前缀和(parallel scan)
    # 利用 GPU 的 Tensor Core 加速
    
    def selective_scan(u, delta, A, B, C, D):
        """
        u: 输入序列 [B, L, d]
        delta: 离散化步长 [B, L, d]
        A: 状态矩阵 [d, N]
        B, C: 投影 [B, L, N]
        D: 跳跃连接 [d]
        """
        # 预计算每个时间步的 effective A
        A_bar = torch.exp(delta.unsqueeze(-1) * A)  # [B, L, d, N]
        B_bar = delta.unsqueeze(-1) * B.unsqueeze(2)  # [B, L, d, N]
        
        # 并行前缀和: 将递推转化为并行归约
        # 使用 Blelloch 扫描算法,O(log n) 的并行步骤
        h = parallel_scan(A_bar, B_bar * u.unsqueeze(-1))
        
        # 输出投影
        y = (h @ C.unsqueeze(-1)).squeeze(-1) + D * u
        return y

    这个算法的精妙之处在于:它牺牲了一点数学上的等价性(用并行前缀和替代严格递推),换取了 GPU 上的高效并行。现代 GPU 的 Tensor Core 可以极大加速这个计算流程。


    四、架构对比:为什么 Mamba 能行

    4.1 与 Transformer 的核心差异

    时间复杂度(训练) O(n²·d) O(n·d²) 时间复杂度(推理) O(n·d)(单步) O(d)(每步固定计算) 内存消耗(训练) O(n² + n·d) O(n·d) 长距离依赖 天然支持(任意距离) 通过全局状态隐式建模 Inductive Bias 几乎无(完全数据驱动) 受控系统的稳定性先验
    维度 Transformer Mamba

    4.2 理论分析:选择性信息的压缩视角

    从信息论角度看,Mamba 做的事情是:将整个序列选择性地压缩进一个固定大小的隐状态中。这类似于无损压缩 vs 有损压缩的权衡:

    • Transformer 的完整 KV Cache 可以看作是无损存储——完整保留了所有历史信息
    • Mamba 的隐状态是有损压缩——只保留对后续预测"有用"的信息

    关键问题是:这种有损压缩的质量如何?Mamba 论文的实验表明,在 选择性任务(selective tasks)——如选择性复制(Selective Copying)、诱导头(Induction Heads)等——上,Mamba 能完美完成,因为它学会了丢弃无关信息、保留关键信息。但在需要精确 recall 的任务上,Mamba 仍有差距。


    五、代码实战:从零实现选择性扫描

    让我们从底层实现一个选择性扫描(Selective Scan)模块:

    import torch
    import torch.nn as nn
    import torch.nn.functional as F
    from einops import rearrange
    
    class SelectiveScan(nn.Module):
        """简化的选择性扫描实现"""
        
        def __init__(self, d_model, d_state=16):
            super().__init__()
            self.d_model = d_model
            self.d_state = d_state
            
            # A 矩阵:初始化为 S4 论文建议的特征值分布
            # 这里简化为可学习的下三角矩阵
            self.A_log = nn.Parameter(torch.log(
                torch.arange(1, d_state + 1, dtype=torch.float32).unsqueeze(0).repeat(d_model, 1)
            ))
            
            # 投影层
            self.x_proj = nn.Linear(d_model, d_state * 3, bias=False)  # delta, B, C
            self.dt_proj = nn.Linear(d_state, d_model, bias=True)
            self.out_proj = nn.Linear(d_model, d_model, bias=False)
            
            # 初始化 dt_proj 的 bias,使 delta 初始值较小
            nn.init.uniform_(self.dt_proj.bias, -0.5, 0.5)
            
            self.D = nn.Parameter(torch.randn(d_model))
        
        def forward(self, u):
            """
            u: [batch, seq_len, d_model]
            returns: [batch, seq_len, d_model]
            """
            batch, seq_len, d_model = u.shape
            
            # 计算 A(保证稳定性)
            A = -torch.exp(self.A_log)  # [d_model, d_state]
            
            # 投影得到 delta, B, C
            x_proj_out = self.x_proj(u)  # [batch, seq_len, d_state * 3]
            delta, B, C = x_proj_out.chunk(3, dim=-1)
            
            # delta 通过 softplus 确保为正
            delta = F.softplus(self.dt_proj(delta))  # [batch, seq_len, d_model]
            
            # 执行选择性扫描(简化版递推实现)
            # 实际生产中应使用 CUDA 优化的 parallel scan
            h = torch.zeros(batch, d_model, self.d_state, device=u.device)
            ys = []
            
            for t in range(seq_len):
                # 更新规则:h_t = exp(-delta * A) * h_{t-1} + delta * B * x
                dA = torch.exp(-delta[:, t].unsqueeze(-1) * A.unsqueeze(0))  # [batch, d_model, d_state]
                dB = delta[:, t].unsqueeze(-1) * B[:, t].unsqueeze(1)       # [batch, d_model, d_state]
                h = dA * h + dB * u[:, t].unsqueeze(-1)
                
                # 输出:y = C @ h
                y_t = (h * C[:, t].unsqueeze(1)).sum(dim=-1)  # [batch, d_model]
                ys.append(y_t)
            
            y = torch.stack(ys, dim=1)  # [batch, seq_len, d_model]
            
            # 跳跃连接
            y = y + self.D.unsqueeze(0).unsqueeze(0) * u
            
            return self.out_proj(y)
    
    
    class MambaBlock(nn.Module):
        """完整的 Mamba 块:SSM + MLP"""
        
        def __init__(self, d_model, d_state=16, d_conv=4, expand=2):
            super().__init__()
            self.d_inner = d_model * expand
            
            # 输入投影
            self.in_proj = nn.Linear(d_model, self.d_inner * 2, bias=False)
            
            # 卷积层(提取局部特征)
            self.conv1d = nn.Conv1d(
                self.d_inner, self.d_inner, 
                kernel_size=d_conv, 
                padding=d_conv-1, 
                groups=self.d_inner
            )
            
            # 选择性扫描
            self.ssm = SelectiveScan(self.d_inner, d_state)
            
            # 门控
            self.act = nn.SiLU()
            self.out_proj = nn.Linear(self.d_inner, d_model, bias=False)
        
        def forward(self, x):
            """
            x: [batch, seq_len, d_model]
            """
            # 残差分支
            x_res = x
            xz = self.in_proj(x)
            x, z = xz.chunk(2, dim=-1)
            
            # 卷积提取局部特征
            x = rearrange(x, 'b l d -> b d l')
            x = self.conv1d(x)[..., :x.shape[-1]]  # 截断 padding
            x = rearrange(x, 'b d l -> b l d')
            x = self.act(x)
            
            # 选择性扫描
            x = self.ssm(x)
            
            # 门控
            z = self.act(z)
            x = x * z
            
            return self.out_proj(x) + x_res
    
    
    class SimpleMamba(nn.Module):
        """一个极简但完整的 Mamba 语言模型"""
        
        def __init__(self, vocab_size, d_model=256, n_layers=6, d_state=16, max_len=2048):
            super().__init__()
            self.embedding = nn.Embedding(vocab_size, d_model)
            self.pos_emb = nn.Embedding(max_len, d_model)
            
            self.layers = nn.ModuleList([
                MambaBlock(d_model, d_state) for _ in range(n_layers)
            ])
            
            self.norm = nn.RMSNorm(d_model)
            self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
        
        def forward(self, input_ids):
            B, L = input_ids.shape
            pos = torch.arange(L, device=input_ids.device).unsqueeze(0).expand(B, -1)
            
            x = self.embedding(input_ids) + self.pos_emb(pos)
            
            for layer in self.layers:
                x = layer(x)
            
            x = self.norm(x)
            logits = self.lm_head(x)
            return logits

    六、工程实践中的关键挑战

    6.1 训练不稳定性

    Mamba 训练的稳定性比 Transformer 更具挑战性。选择性步长 Δ 如果训练不当,会导致状态爆炸或消失。实践中需要注意:

    # 训练技巧 1:梯度裁剪必须更激进
    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.1)
    
    # 训练技巧 2:Δ 的初始化非常关键
    # 初始时应让 delta 较小(接近 0.1),避免训练初期状态剧烈变化
    # 论文建议 dt_proj 的 bias 初始化在 [-0.5, 0.5]
    
    # 训练技巧 3:A 矩阵的特征值分布
    # 建议使用 HiPPO 初始化或类似的连续时间先验
    # 确保特征值的负实部足够大,保证系统稳定性

    6.2 长序列中的遗忘问题

    尽管理论上 SSM 可以建模无限长序列的依赖,但在实践中,当序列超过一定长度后,早期信息会被逐渐稀释。这是因为选择性步长 Δ 的累积效应——如果后面不断有新信息输入,早期状态的对数概率贡献会指数衰减。

    实验表明,Mamba 在 DNA 序列建模(需要 64K+ 上下文)上表现出色,但在需要精确回忆长距离 token 的任务上,仍不如 Transformer。

    6.3 推理引擎的优化

    Mamba 的推理有一个独特优势:每一步的计算量是固定的(不像 Transformer 需要访问整个 KV Cache)。这使得它在以下场景有显著优势:

    • 实时对话系统:生成延迟稳定,无随对话长度增长而变慢的问题
    • 流式处理:可以无限处理输入,内存占用恒定
    • 批处理推理:吞吐量不会因为最长序列的存在而受到严重影响
    # Mamba 推理伪代码:每步固定计算量
    def mamba_generate_step(model, token, state):
        """每次调用只执行固定计算,不访问历史"""
        x = model.embedding(token)  # [B, 1, d]
        for layer in model.layers:
            # 状态更新:h_new = f(h_old, x)
            state[layer.idx], x = layer.sigmoid_ssm(x, state[layer.idx])
        logits = model.lm_head(x)
        return logits, state  # 返回新状态,下一步使用

    七、混合架构:SSM 与 Attention 的联姻

    2024 年下半年以来,越来越多的模型开始探索混合架构——将 SSM 的高效和 Attention 的精确结合:

    7.1 Jamba 架构(AI21 Labs)

    Jamba 采用 8:1 的 Mamba-to-Attention 层比例,在保持接近纯 Transformer 性能的同时,将内存占用降低了 4-8 倍:

    class JambaLayer(nn.Module):
        def __init__(self, config, use_attention=False):
            super().__init__()
            if use_attention:
                self.layer = AttentionBlock(config)
            else:
                self.layer = MambaBlock(config)
        
        def forward(self, x):
            return self.layer(x)
    
    # 配置:每 8 层 Mamba 加 1 层 Attention
    jamba_config = {
        'total_layers': 52,
        'attention_every_k': 8,  # 第 8, 16, 24... 层使用 Attention
        'experts_per_layer': 4   # MoE 化的中间层
    }

    7.2 Zamba(Zyphra)

    Zamba 更进一步,仅用 1 层 Attention 就恢复了完整的 Transformer 性能。其核心研究发现:Attention 的价值不在于数量,而在于位置——在信息流的"瓶颈"处(如跨段落边界)放置 Attention 层,效果远胜于均匀分布。

    7.3 工程启示

    这对架构设计的启示是:

    1. SSM 适合处理"连续流"——正文、对话等自然序列
    2. Attention 适合处理"结构跳变"——段落开头、列表项等边界
    3. 比例比类型更重要——适量 Attention + 大量 SSM 是当前最佳实践

    4. 八、生产部署考量

      8.1 推理框架支持

      截至 2026 年,Mamba 的推理生态已日趋成熟:

      vLLM 部分 主流 LLM 推理引擎,SSM 支持正在合并 Transformers 完整 HuggingFace 官方支持 Mamba/Mamba2 mistral.rs 完整 Rust 实现,高性能 SSM 推理 MambaCUDA 完整 官方 CUDA kernel,最优性能 Tinygrad 轻量 适合边缘部署,代码可读性高
      框架 SSM 支持 特性

      8.2 内存与吞吐对比

      以 Llama-3 70B 大小的模型为例,处理 32K 上下文时:

      • 纯 Transformer:KV Cache ~ 128GB,单步延迟 ~ 45ms
      • 纯 Mamba:状态 ~ 16GB(固定),单步延迟 ~ 8ms
      • 混合架构(20:1):状态 ~ 20GB + KV Cache ~ 6.4GB,单步延迟 ~ 12ms

      在长上下文场景下,Mamba 的优势不仅是成本,更是可行性——有些场景对 Transformer 而言根本不可行(如整本书籍的实时分析),Mamba 让这些场景成为可能。


      九、未来方向

      SSM 领域仍在快速演进。以下几个方向值得关注:

      9.1 Mamba-2:对角化与tensor并行

      Mamba-2 将 SSM 参数化为对角矩阵(Diagonal SSM),使得矩阵幂运算从 O(N²) 降至 O(N)。更关键的是,对角化让张量并行变得可以直接实现:

      # Mamba-2 的对角 SSM:可以按状态维度切分到不同 GPU
      # 每个 GPU 负责 d_state / world_size 个状态维度
      # 通信开销远小于 Attention 的 KV Cache 同步
      class DiagonalSSM(nn.Module):
          def __init__(self, d_model, d_state):
              # 只需要存储对角元素,参数减少 N 倍
              self.Lambda = nn.Parameter(torch.randn(d_model, d_state))  # 对角矩阵
          
          def forward(self, u):
              # Lambda^k 只需逐元素幂运算,无需矩阵乘法
              Lambda_power = self.Lambda.unsqueeze(0).unsqueeze(0) ** torch.arange(seq_len)
              # ... 向量化实现

      9.2 结构化稀疏 SSM

      将 SSM 的 A 矩阵从稠密改为结构化稀疏(块对角、带状),可以在保持长距离建模能力的同时,进一步降低计算开销。

      9.3 硬件亲和设计

      Mamba 的固定计算图特性天然适合专用 AI 加速器。与 Attention 需要不同形状的矩阵乘加不同,SMA 每一步的计算模式完全相同,这为 FPGA/ASIC 实现提供了天然优势。


      结语

      Mamba 和状态空间模型代表了一个重要趋势:当我们从第一性原理审视序列建模问题时,会发现 Transformer 并非唯一的答案,甚至不是最优答案。

      选择性状态空间的核心贡献不在于它"替代"了 Transformer,而在于它证明了另一条路径的存在——一条路径可以用线性复杂度处理序列,用固定内存承载长距离依赖,同时保持与 Transformer 相媲美的性能。

      对于工程师而言,2026 年的技术版图已经清晰:Transformer 不会消失,但它不再是唯一的选择。理解 SSM、理解选择性扫描、理解混合架构的设计权衡,已经成为现代 AI 基础设施工程师的核心素养之一。


      参考论文:

      • Gu, A., & Dao, T. (2023). Mamba: Linear-Time Sequence Modeling with Selective State Spaces. arXiv:2312.00752
      • Gu, A., et al. (2022). Efficiently Modeling Long Sequences with Structured State Spaces (S4). ICLR 2022
      • Dao, T., & Gu, A. (2024). Mamba2: Subquadratic-State Space Models Are As Efficient As Transformers. arXiv:2405.21060
      • Lieber, O., et al. (2024). Jamba: A Hybrid Transformer-Mamba Language Model. arXiv:2408.15668
点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部