vLLM 中用于混合 SSM 模型的解耦推理服务

15 分钟阅读
Nicolò Lucchesi, Zhanqiu Hu (Red Hat) 以及 vLLM 团队

简介

将 Mamba 风格的 SSM 层与标准全注意力 (FA) 层交错的混合架构(例如 NVIDIA Nemotron-H)正日益受到关注,因为它们结合了状态空间模型的线性时间效率和注意力机制的表达能力。vLLM 已经通过其 基于 NIXL 的 KV 连接器 支持标准 Transformer 模型的分离式预填充/解码 (P/D):预填充实例计算 KV 缓存块,而解码实例通过 RDMA 拉取这些块,从而消除了冗余的重新计算。但将此扩展到混合模型并非易事。FA 和 SSM 层存储的状态在本质上是不同的,布局和大小也各异,而块管理器和 NIXL 连接器最初是围绕单一、统一的 KV 缓存格式设计的。

在这篇文章中,我们描述了如何扩展 NIXL 连接器以支持分离模式下的混合 SSM-FA 模型。关键思想包括:

  • 双描述符视图 — 两套 NIXL 块描述符,它们以不同的偏移量和大小索引相同的物理内存区域,一套用于 FA 块,另一套用于 SSM 块。
  • 物理/逻辑块桥接 — 处理块管理器所见的逻辑块抽象与注意力内核所需的物理块大小之间的不匹配。
  • 三描述符卷积传输 — 对 Mamba 卷积状态进行分解,从而在发送端无需重新排列数据即可实现异构张量并行传输。

这些更改都不会修改标准 Transformer 模型现有的工作流。它们是纯粹的增量式扩展,仅在模型包含 SSM 层时激活。此功能在 vllm>=v0.20.0 中可用。

这项工作基于 NIXL 的 HMA 接口,并跨越了多个 PR:

  • #36687 — 用于混合 SSM-FA 模型的双描述符视图和同构 TP 支持
  • #37416 — 用于 Mamba 内核的 DS 卷积状态布局
  • #37635 — 异构 TP 三描述符卷积状态传输
  • #37310 — 用于 Mamba P/D 分离的 N-1 预填充

背景:NIXL KV 传输工作流

在深入探讨混合模型更改之前,让我们简要回顾一下 NIXL 分离式 P/D 如何为标准 Transformer 工作。

工作流分为四个阶段:

  1. 注册内存区域 — 每个 worker 向 NIXL 注册其 KV 缓存张量,以便可以通过 RDMA 访问它们。
  2. 创建块描述符 — 对于每个注册的区域,我们创建每个块的描述符,指定 (address, length, device_id)。这些描述符是我们传输的单位:我们传输的是单个块,而不是整个区域。
  3. 握手 — 当解码 (D) worker 首次需要从预填充 (P) worker 拉取数据时,两者交换元数据:代理句柄、块计数、块长度等。每个 P-D 对执行一次。
  4. 传输 — 调度程序告诉 D 要从 P 拉取哪些块。D 将 block_id -> descriptor_id 进行映射,发起 RDMA READ,并轮询完成情况。

对于具有 M 个注册区域和 N 个块的标准模型,描述符列表如下所示:

+----------------------------------+
| Region 0: desc_0 ... desc_{N-1}  |
| Region 1: desc_0 ... desc_{N-1}  |
| ...                              |
| Region M: desc_0 ... desc_{N-1}  |
+----------------------------------+

区域 r 中的块 ID b 映射到描述符索引 r * N + b

混合模型的挑战在于这种统一方案不再适用:FA 层和 SSM 层需要不同的描述符大小和块计数。


挑战:FA 和 SSM 状态在本质上是不同的

在标准 Transformer 中,每一层的 KV 缓存具有相同的形状:[num_blocks, 2, block_size, num_kv_heads, head_dim](或类似的布局变体)。所有层共享相同的块大小、页面大小和块数。

Mamba 层存储的内容完全不同。它们不存储每个 Token 的 K/V 对,而是维护一个折叠的卷积状态时间 SSM 状态

Conv state:  (conv_dim, state_len)    e.g. (3072, 3)   -- bf16
SSM state:   (num_heads, head_dim, state_size)   e.g. (32, 64, 128) -- fp32

这些状态中没有“Token”的概念 — 它们是整个序列历史的固定大小摘要。这意味着 SSM 的 block_size 实际上为 1:每个块是一个完整的状态快照,而不是一组每个 Token 的向量。请记住:块是此处的单个传输单位

HMA 共享张量布局

vLLM 的混合内存分配器 (HMA) 按类型对层进行分组:所有 FA 层为一组,所有 SSM 层为另一组,依此类推。然后,它跨组进行内存池化,使得每个组中相同位置的层共享同一个物理张量。这样做非常高效(块是可互换的),但这意味着同一个张量被一组视为 FA 块,同时被另一组视为 SSM 块。

这是 Nemotron-H 等模型的布局结果:

                KV Cache Tensor (shared via HMA pooling)
                 /                        \
                /                          \
     Attention (FA) View              Mamba View
              |                            |
    +-----------------------+    +-----------------------+
    | Block 0               |    | Block 0               |
    |   Key     |  Value    |    |  Conv |    SSM  |[pad]|
    | Block 1               |    | Block 1               |
    |   Key     |  Value    |    |  Conv |    SSM  |[pad]|
    |  ...                  |    |  ...                  |
    +-----------------------+    +-----------------------+

页面大小不同:FA 页面由 block_size * num_kv_heads * head_dim 控制(K/V 乘以 2),而 SSM 页面为 conv_state_bytes + ssm_state_bytes。HMA 会增加 FA 的 block_size 直到它大于 Mamba 的块大小,然后对 Mamba 行进行填充(+[pad]),从而使两组在字节数上具有相同的页面大小,实现了共享张量方案。

NIXL 面临的问题:单个具有统一 (address, length) 条目的描述符列表无法正确索引两个视图。我们需要在单独的描述符上注册 K/V(以及类似的 Conv/SSM),以便在异构设置(即 D TP != P TP 时)下对 K/V 头进行索引。

b 的 FA 描述符指向 base + b * page_size,长度为 fa_block_len。同一块 b 的 Mamba 描述符指向相同的 base + b * page_size,长度为 conv_sizessm_size。这些是不同的。


双描述符视图

我们的解决方案是在同一块物理内存上注册两个独立的描述符列表,并将它们连接起来,由单个 NIXL 传输句柄指向:

+------------------------------------------------------+
|  FA descriptors (M regions x N_phys blocks)          |
|                                                      |
|  Region 0                                            |
|    FA_desc_K[0], FA_desc_K[1], ... FA_desc_K[N-1]    |
|    FA_desc_V[0], FA_desc_V[1], ... FA_desc_V[N-1]    |
|  Region 1                                            |
|    ...                                               |
|  Region M                                            |
|    ...                                               |
|                                                      |   ^
|  --------------------------------------------------- |   | num_descs
|                                                      |   v
|  Mamba descriptors (M regions x N_log blocks)        |
|                                                      |
|  Region 0                                            |
|    Mamba_desc_x[0]   ... Mamba_desc_x[N-1]           |
|    Mamba_desc_B[0]   ... Mamba_desc_B[N-1]           |
|    Mamba_desc_C[0]   ... Mamba_desc_C[N-1]           |
|    Mamba_desc_SSM[0] ... Mamba_desc_SSM[N-1]         |
|  Region 1                                            |
|    ...                                               |
|  Region M                                            |
|    ...                                               |
+------------------------------------------------------+

注意:请记住我们使用 N_phys/_log 分别表示物理块和逻辑块。您可以假设 N_phys=N_log=N,若非如此请参阅下一节。

注意:上面的 Mamba 部分已经反映了卷积状态分解为 x, B, C 子投影,将在下文的 三描述符卷积传输 中解释。对于同构 TP,这些简化为两个子区域(Conv, SSM)。

FA 描述符占据前 num_descs = M * N_phys 个槽位。Mamba 描述符紧随其后。块 ID 映射变为:

if is_fa_group:
    desc_id = region_id * N_phys + block_id
else:  # mamba group
    desc_id = mamba_region_id * N_log + block_id + num_descs

物理块大小与逻辑块大小

第二个复杂之处源于注意力内核的要求。像 FlashInfer 这样的后端需要特定的物理块大小(例如 16 个 Token),这可能与用户设置或 HMA 计算出的逻辑块大小不同。

对于标准模型,这通过一个简单的比率来处理:

physical_blocks = logical_blocks * ratio
ratio = logical_block_size / kernel_block_size

对于混合模型,该比率仅适用于 FA 层。SSM 层没有“Token”维度可供拆分,因此它们总是直接使用 logical_blocks。这意味着描述符列表的 FA 和 Mamba 部分使用不同的块计数:

FA section:    M regions * N_phys blocks    (N_phys = N_logical * ratio)
Mamba section: M regions * N_logical blocks

这通过 _physical_blocks_per_logical 字段进行跟踪,该字段是按引擎计算的(因为当 TP 大小不同时,P 和 D 可能具有不同的比率)。_get_block_descs_ids 中的块 ID 到描述符 ID 映射使用适当的步幅,具体取决于它是解析 FA 组还是 Mamba 组。


三描述符卷积传输

对于同构 TP(P 和 D 使用相同的 --tensor-parallel-size),传输 SSM 状态很简单:每个 D rank 从匹配的 P rank 读取相应的 conv + SSM 块。

异构 TP 使这变得更加困难。考虑 P_TP=1, D_TP=4:四个 D worker 必须各自从单个 P worker 读取它们所分片的 conv 和 SSM 状态。SSM 时间状态沿 heads 维度进行分片,这是第一个轴——所以切片非常简单。但卷积状态的结构如下:

Conv state = [x | B | C]     where x, B, C are sub-projections
              ^   ^   ^
              |   |   |
     intermediate_size / TP   groups_ss / TP   groups_ss / TP

使用标准 SD 布局 (state_len, dim),这些子投影在内存中是交错的。一个只需要部分 x 的 D worker 需要收集不连续的字节——这对于零拷贝 RDMA 是不切实际的。

DS 布局解决方案

我们需要用于卷积状态的 DS 布局 (dim, state_len)(通过 VLLM_SSM_CONV_STATE_LAYOUT=DS 设置)。在这种布局中,每个子投影的数据在内存中是连续的:

DS layout within one page:

|--- x (x_bytes) ---|--- B (b_bytes) ---|--- C (b_bytes) ---|--- SSM ---|

每个 D rank 现在可以通过三次单独的连续 RDMA 读取来读取其 xBC 的切片——这就是所谓的“三描述符传输”(我们仍然只发起一个 NIXL READ)。

对于异构 TP,remote_conv_offsets 方法计算每个 D rank 的切片在 P 页面内的位置,并考虑了 TP 比率。这给了我们每个 Mamba 层 4 个描述符区域(x, B, C, SSM),而不是同构情况下的 2 个区域(Conv, SSM)。代价是更大的描述符列表,但 RDMA 传输本身仍然是高效的连续读取。

没有额外的内存中暂存缓冲区在任何 GPU 上分配。任何一方都不需要进行数据重新排列

注意:我们在常规托管设置中使用 DS 布局时,没有测量到明显的内核性能回归。我们可能会在未来版本中将标准布局始终设为 DS。

零开销:无需额外缓冲区,无需置换

一个更简单的替代方案是向每个 D rank 传输整个卷积状态,然后在本地将其排列/切片成正确的形状。但对于 Mamba,我们刻意避免了这种方法:

  • 无暂存缓冲区 — 在 D 上进行排列需要在每个 D worker 上分配一个 P 的完整卷积状态大小的临时缓冲区。对于像 Nemotron-H 这样的模型,每个块的卷积状态已经很大(bf16 中为 3 * 3072 * 2 字节)。乘以数千个块和所有 Mamba 层,这种开销会增加,并占用原本可以用于 KV 缓存的空间。
  • 无传输后重新排列 — 使用 DS 布局,每个 D rank 只读取它需要的确切字节,直接进入 KV 缓存中的最终目的地。没有传输后的内核来重新排列数据。传输完成,状态即可立即使用。
  • 只传输你拥有的 — 每个 D rank 只传输其 1/TP 的卷积状态份额,而不是完整状态。对于 D_TP=4,这意味着与“传输所有内容,本地切片”的方法相比,每个 rank 的数据量减少了 4 倍。
  • 跳过 HMA 填充 — 回想一下 HMA 填充了 SSM 页面以使其匹配 FA 页面大小。Mamba 描述符的大小调整为实际的 conv_bytes + ssm_bytes,而不是填充后的页面大小。这意味着我们永远不会通过线路传输填充字节——只有真实状态。对于填充量很大的模型(例如,当 FA 页面大小远大于原始 SSM 状态时),这可以显著减少每个块的传输量。

下图验证了 Nemotron Super 120B 在 TP=4(FA block_size=4224,由 HMA 设置)下的零开销传输优化。对于每个 KV 缓存数据类型(bf16 和 fp8),我们将朴素基准——它为 Mamba 块传输完整的 HMA 填充页面——与最优方法进行了比较,后者只传输实际的 conv + SSM 字节,跳过所有 HMA 填充和/或辅助缓冲区。我们首先验证我们的方法是否如传输指标所报告的那样匹配最优

对于 fp8,FA 页面大小较小(每个元素 1 字节对比 2 字节),因此在此配置中填充可忽略不计。然后我们展示了在 bf16 设置下的节省,我们的方法消除了每个请求约 50 MB 的不必要传输。

由于 Mamba 状态是每个请求的固定大小摘要,因此随着 ISL 的增加,传输大小会随 FA 块的数量而扩展。

Figure 1: P→D transfer volume vs. input sequence length for Nemotron Super 120B (TP=4, FA block_size=4224). The Naive and Optimal baselines are computed analytically from the model's page sizes and block counts. The Measured line reports the actual bytes transferred (as reported by NIXL) during disaggregated P/D serving. Our approach (Optimal) eliminates HMA padding overhead, which is reflected in the measured transfer.
图 1:Nemotron Super 120B 的 P→D 传输量与输入序列长度的关系(TP=4, FA block_size=4224)。朴素基准和最优基准是根据模型的页面大小和块计数解析计算的。测量线报告了在分离式 P/D 服务期间(由 NIXL 报告)实际传输的字节数。我们的方法(最优)消除了 HMA 填充开销,这反映在测量到的传输中。

整合:Nemotron-H 示例

让我们通过一个具体的例子:在 TP=2 下以分离式 P/D 服务 nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-FP8

模型结构:总共 52 层,在 Mamba 和 FA 之间交替。HMA 将它们分为 5 组(4 个 Mamba,1 个 FA)。经过内存池化后,产生了 6 个共享的 KV 缓存张量。

KV 缓存布局:

FA layers:    [num_blocks, 2, block_size=400, 4, 128]   # K/V with HMA-inflated block_size
SSM layers:   [num_blocks, 3, 3072]  (conv)  +  [num_blocks, 48, 64, 128]  (ssm)

HMA 填充块大小,以便两个视图在字节数上具有相同的页面大小。内核(FlashInfer/FlashAttention)可能会进一步细分 FA 块,创建一个物理/逻辑比率。

描述符注册:

  1. 这 6 个共享张量被注册为 NIXL 内存区域(与密集模型相同)。
  2. 为所有 6 个区域 x N_phys 块创建 FA 描述符,分别索引 K 和 V。
  3. 附加 Mamba 描述符:6 个区域 x N_logical 块,每个有 4 个子区域(x, B, C, SSM)用于三描述符传输。

传输流:

  1. P 完成预填充。调度程序按组分配块 ID:[[fa_block_ids], [mamba_block_ids_g0], [mamba_block_ids_g1], ...]
  2. D 接收块 ID 并将其映射到描述符索引:FA 块使用标准的 region * N + block_id 公式;Mamba 块添加 num_descs 偏移量并使用 N_logical 步幅。
  3. D 发起单个 make_prepped_xfer READ,同时使用 FA 和 Mamba 描述符,然后轮询完成情况。
  4. 完成后,D 通知 P,以便 P 可以释放这些块。

从 D 的角度来看,整个传输是一个单一的异步操作。没有中间缓冲区,没有数据重排。


性能

我们将分离式 P/D 与在通过 NVLink 连接的 8x H200 GPU 上进行的托管式服务进行了基准测试。该模型是 nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-FP8,这是一个近期的 120B LatentMoE 混合架构,具有交错的 Mamba2 和全注意力层。

  • 托管基准:单个实例,TP=8,全部 8 个 GPU。
  • 分离式 P/D:1 个预填充实例(TP=4, 4 个 GPU)+ 1 个解码实例(TP=4, 4 个 GPU),GPU 总数相同。

我们将并发从 8 扫到 256 个并发用户,并绘制每个 GPU 的输出吞吐量与每个用户的输出 Token 率(交互性)的关系图。工作负载使用 ShareGPT 作为测试数据集。

所有运行都使用非常高的预热值,以确保 KV 缓存被“打乱”,从而避免请求块刚好连续分配时所带来的性能提升。这更准确地反映了常规的长期运行使用。还可以通过检查指标中报告的恒定数量的描述符(在整个数据集扫描过程中)来验证这一点。

Figure 2: Disaggregated P/D vs. co-located serving for a hybrid SSM model. Throughput-vs-latency Pareto curve across concurrency levels. Prefix-caching disabled.
图 2:混合 SSM 模型的分离式 P/D 与托管式服务对比。各并发水平下的吞吐量-延迟帕累托曲线。前缀缓存已禁用。

结果显示了与标准 Transformer 模型的分离式服务所观察到的相同模式:分离式 P/D 在更高的批处理大小时帕累托占优于托管基准。通过将解码与预填充干扰隔离,解码实例可以在不中断的情况下维持更大的批处理,从而在高并发下每个 GPU 产生显著更高的输出 Token/s。


入门指南

要运行带有分离式 P/D 的混合 SSM 模型:

# Prefill instance
VLLM_SSM_CONV_STATE_LAYOUT=DS vllm serve nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-FP8 \
    --tensor-parallel-size 2 \
    --gpu-memory-utilization 0.85 \
    --trust-remote-code \
    --max-model-len 8192 \
    --block-size 128 \
    --no-disable-hybrid-kv-cache-manager \
    --kv-transfer-config '{"kv_connector":"NixlConnector","kv_role":"kv_both"}'

注意:异构 TP 需要设置 DS 卷积状态布局 VLLM_SSM_CONV_STATE_LAYOUT=DS,否则无需设置。


局限性与未来工作

  • Mamba1 模型:三描述符卷积传输目前仅支持 Mamba2。Mamba1 的 SSM 时间形状 (intermediate_size // tp, state_size) 不允许重构卷积分解所需的 intermediate_size。同样,GDN 支持(Qwen3.5+)已列入分离式 路线图
  • 投机解码:SSM 状态传输与投机解码之间的交互尚未经过广泛验证。
  • HMA 混合块大小:启用 HMA 时,尚不支持 P 和 D 之间的不同块大小(block_size_ratio > 1)。

致谢

Thomas Parnell (IBM Research), Roi Koren (NVIDIA)