VarLen:变长序列推理如何消灭Padding

VarLen 用「铺平 + cu_seqlens 索引」的数据布局替代等长 padding,把 batch 里长短不一的序列压进一维连续缓冲区,显存、带宽、计算三重浪费一起省掉。从 cu_seqlens 数据结构,到 FlashAttention 的 varlen 接口,再到训练侧 sample packing 的姿势,最终落到能力边界——逻辑浪费归 VarLen 管,硬件粒度碎片归 GPU SIMT执行模型:SM、Warp与算力碎片 管。

Padding 的账:浪费到底有多大

LLM 推理的请求长度天然参差:系统 prompt 几千 token,用户问题可能只有几十。但 batch 维的算子要求规整张量,传统做法是把整个 batch padding 到最长序列:

batch = [100, 500, 1000] token,padding 到 1000:
 
s1  ████████░░░░░░░░░░░  100 有效 + 900 无效
s2  ██████████████░░░░░  500 有效 + 500 无效
s3  ███████████████████  1000 有效
 
形状 (3, 1000) = 3000 token,有效仅 1600,利用率 53%

浪费是三重的:显存按 padded 长度分配激活和 KV cache;带宽搬运的是无效 token;计算——attention 对 padding 位置照算不误。batch 越大、长度方差越大,这笔账越触目惊心。

核心设计:铺平 + cu_seqlens

VarLen 的思路非常直接:既然浪费来自「补齐」,那就不补。把所有序列首尾相接铺平成一维缓冲区,再用一个前缀和数组 cu_seqlens 记录每条序列的起止位置:

铺平后 (1600,):
 
s1 ████████│s2 ██████████████│s3 ███████████████████
0        100              600                   1600
 
cu_seqlens = [0, 100, 600, 1600]
第 i 条序列占据铺平缓冲区的 [cu_seqlens[i], cu_seqlens[i+1])

kernel 内部按 cu_seqlens 切段,每段独立计算自己的 attention,段与段之间互不可见。缓冲区里没有一个无效 token,三重浪费同时消失。这也是 FlashAttention、vLLM 等现代推理栈的底层布局——不 padding,直接铺。

实战用法

flash_attn_varlen_func

FlashAttention 2.x 的变长入口。q/k/v 不再是 (batch, seqlen, heads, dim),而是 (total_tokens, heads, dim)

import torch
from flash_attn import flash_attn_varlen_func
 
seqlens = torch.tensor([100, 500, 1000], dtype=torch.int32, device="cuda")
total = int(seqlens.sum())  # 1600,而非 padding 后的 3000
 
# 前缀和并在头部补 0:[0, 100, 600, 1600]
cu_seqlens = torch.nn.functional.pad(seqlens.cumsum(dim=0), (1, 0))
 
heads, dim = 32, 128
q = torch.randn(total, heads, dim, dtype=torch.bfloat16, device="cuda")
k = torch.randn_like(q)
v = torch.randn_like(q)
 
out = flash_attn_varlen_func(
    q, k, v,
    cu_seqlens_q=cu_seqlens,
    cu_seqlens_k=cu_seqlens,  # 自注意力:q/k 同源,边界一致
    max_seqlen_q=1000,        # 只影响 kernel launch 配置,取 batch 内最大值
    max_seqlen_k=1000,
)
# out.shape == (1600, heads, dim),逐段注意力,无 padding 参与计算

两个细节容易踩坑:cu_seqlens 必须是 int32 且在 GPU 上;max_seqlen 传大了不会算错,但会浪费 launch 资源,传小了直接越界。

cu_seqlens 构造一行流

F.pad(x.cumsum(0), (1, 0)) 是标准写法:cumsum 得到右端点,左侧补 0 凑出 batch+1 个边界。自己手写循环拼接容易在空序列(长度 0)上出 off-by-one。

训练侧:Sample Packing

训练场景同样受益。把多个短样本拼满一个 4096 的长序列再组 batch,GPU 上的有效 token 密度大幅提升——前提是注意力不能跨界:文档 A 的 token 不该 attend 到文档 B。做法完全一样,拼好的序列配一个 cu_seqlens,FA2 在段边界处截断注意力。预训练语料喂给开源框架(如 Megatron)时,sample packing + varlen attention 已是标准配置。

生态位

推理框架不直接暴露这个 API,但思想同源:vLLM、SGLang 的 continuous batching 每个 step 动态重组 batch,配合 PagedAttention 管理 KV cache,本质上都是在「消灭 padding」这条线上做文章,见 vLLM使用指南:大模型高吞吐推理的事实标准SGLang:大模型结构化生成语言与推理框架。PyTorch 官方的 NestedTensor 也在补原生变长支持,成熟度尚不如显式 cu_seqlens 方案。

适用场景与边界

场景VarLen 收益原因
Prefill,序列长短差异大padding 浪费占比大,铺平直接全免
训练 sample packing提升 token 密度,同 batch 吞吐显著上升
小 batch 短序列推理显存/带宽收益保留,但算力利用率天花板卡在 Warp 粒度
Decode 阶段每步每序列只产 1 token,天然近似等长

边界在哪?VarLen 消灭的是逻辑浪费——缓冲区里不再有无效 token。但 GPU 硬件按 32 线程一个 Warp 调度:batch 里只有 17 个有效 token 时,硬件照样分配一个完整 Warp,剩下 15 个线程全程掩码空转。这部分粒度碎片是 SIMT 执行模型的天然属性,任何软件布局都消除不了,只能靠加大 batch 稀释、靠多 Stream 填充空闲 SM——完整分析见 GPU SIMT执行模型:SM、Warp与算力碎片

一句话总结选型:长度方差大、追求吞吐的批式场景,VarLen 几乎总是对的;指望它单独把小 batch 短序列的 MFU 从三成拉到九成,方向就找错了——那得去硬件调度模型里找答案,参见 LLM推理性能终极标尺:MFU 深度解析、计算方法与2nd Forward FLOPs全拆解Prefill 阶段