Mamba

定义

当前 token 主要依赖一个递归更新的 hidden state

经典SSM(Selective State Space Model)

状态转移:‭

$h_t = A h_{t-1} + B x_t$

‬‭‬‭‬输出计算:‭

$y_t = C h_t + D x_t$

‬‭‬‭‬A B C 固定不变,作为卷积训练很快

Mamba

B变成动态,可以根据token的重要程度调整,本质上就是之前的固定的网络需要和输入计算

\[x' = \text{Linear}_x(X)‬\] \[x = \text{SiLU}(\text{Conv1d}(x'))‬‭‬‭‬‭‬‭‬‭‬‭‬‭‬\] \[z = \text{Linear}_z(X)‭\] \[\Delta_t = \text{Softplus}(W_\Delta x_t)\] \[B_t = W_B x_t, \quad C_t = W_C x_t‬‭‬‭‬ ‭‬‭‬‭‬‭\] \[h_t = \exp(\Delta_t A) h_{t-1} + \Delta_t B_t x_t‬‭‬‭‬‭‬‭‬‭ ‭\] \[y_t = C_t h_t‬‭‬‭\] \[\hat{y}_t = y_t \odot \text{SiLU}(z_t)‬‭‬‭‬‭‬‭‬‭‬ ‭\] \[O = \text{Linear}_{\text{out}}(\hat{y})‬‭‬‭‬‭‬‭‬\]

Mamba 和 Attention 对齐

\[Y = \text{Softmax}\left(\frac{Q K^T}{\sqrt{d}}\right) V‬‭\] \[y_i = \sum_{j=1}^i \left( C_i \prod_{k=j+1}^i \bar{A}_k \bar{B}_j \right) x_j\]

如果令构建矩阵 ‭$M_{i,j} = C_i \left( \prod_{k=j+1}^i \bar{A}_k \right) \bar{B}_j$‬‭‬‭‬‭‬‭‬,则:

\[Y = (M \odot L) X\]

其中 ‭$L$‬ 为下三角因果掩码(Causal Mask)。

‬‭‬‭‬‭‬‭‬‭‬‭‬

‬‭‬‭‬‭‬‭‬‭‬‭‬

数据结构

conv state: 短期的滑动窗口,短期的滑动窗口,进行卷积,提取局部上下文,充当平滑器

recurrent / SSM state: 长期递归状态,$h_t = \bar{A}t h{t-1} + \bar{B}_t x_t$ ,负责长距离和选择,推理时存储

State-less / Instantaneous Component

  1. 投影网络

  2. 门控网络 SwiGLU

  3. $\Delta, B, C$ 生成网络

MambaSpec

Code

# vllm/v1/kv_cache_interface.py:689-739

@dataclass(frozen=True)
  class MambaSpec(KVCacheSpec):
      shapes: tuple[tuple[int, ...], ...]
      dtypes: tuple[torch.dtype]
      page_size_padded: int | None = None
      mamba_type: MambaAttentionBackendEnum =
MambaAttentionBackendEnum.MAMBA2
      mamba_cache_mode: str = "none"
      num_speculative_blocks: int = 0

含义

shapes

例子

Mamba1:
 conv_state_shape = (D, kernel_size-1)
 recurrent_state_shape = (D, N)

Mamba2:
 conv_state_shape = (D, kernel_size-1)
 recurrent_state_shape = (num_heads/TP, head_dim, state_size)

GDN:
 conv_state_shape = (D, kernel_size-1)
 recurrent_state_shape = (num_v_heads/TP, head_v_dim, head_k_dim)

KDA:
 conv_state_shape = (D, kernel_size-1)
 recurrent_state_shape = (num_heads/TP, head_dim, head_dim)

dtype

Mamba1/Mamba2:
      (bf16, bf16)

GDN:
      (bf16, bf16)

KDA:
      (bf16, fp32)

mamba_cache_mode

mamba_cache_mode: none: 只保留当前 state align: 保留最近 N 个 state,对齐 prefix cache all: 每个 block 都保存 state

mamba_type

表示这个 state 对应哪种 attention backend。

KVCacheSpec 的子类

整个 vLLM 的 KV cache 系统用统一的 KVCacheSpec 来描述每个 layer 的 cache 布局。

vllm/v1/kv_cache_interface.py:99-130

子类型有:

KVCacheSpec

├── AttentionSpec │ ├── FullAttentionSpec │ ├── SlidingWindowSpec │ ├── MLAAttentionSpec │ ├── CrossAttentionSpec │ ├── EncoderOnlyAttentionSpec │ └── SinkFullAttentionSpec ├── MambaSpec └── …

调度器、cache manager、offloading connector 全部基于这个抽象判断要做什么。

KVCacheSpec 如何创建出

MambaBase