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 阶段。