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
-
投影网络
-
门控网络 SwiGLU
-
$\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 全部基于这个抽象判断要做什么。