声明:本文主要参考开源资料进行学习整理,如有错漏,欢迎评论交流~
引言
1. 为什么 prefill 是显存瓶颈
LLM 推理分两个阶段:
- Prefill:把整条 prompt 一次性喂进模型,计算每个位置的 KV 并写进 cache。这一阶段的注意力是全序列 self-attention(自注意力),prompt 越长,显存峰值越高。
- Decode:一次生成一个 token,每个新 token 只 attend(关注)已有的 KV cache(1 个 query 对 L 个 key)。
问题就在 prefill:一条 1216 token 的 prompt,attention 矩阵是 1216 x 1216 ≈ 148 万个元素,乘以 head 数和 batch,峰值显存很容易顶爆板端 DDR/L2。如果 prompt 再长到几千 token,单次 prefill 直接 OOM。
代价是吞吐:原来 1 次前向完成,现在要 L/c 次。但换来了"能跑"——这是显存受限的端侧场景下的必要取舍。
2. 前置:KV cache 的容量与右对齐

chunk prefill 全程围绕 KV cache 转,先把它的容量和"右对齐"约定讲清楚。
- 容量固定为 C:prefill 的 L 个位置 + decode 生成的新位置,总占用不得超过 C。 L ,且 C - L 是 decode 可生成空间。
- 有效内容右对齐:真实 token 的 K/V 总是写在 cache 的右端(高位)。这源自于“左 padding”——真实 token 在序列右端,进 cache 后自然落在右端。后续 decode 新 token 也是追加在右端往后递增。例如KV cache (长度 C=8,当前有 5 个有效 token):

有这两点先验知识,下面的滚动写入和 mask 切分才有据可依。
3. 核心机制:滚动 KV cache

每个 chunk 算出的 K/V,追加到 KV cache 的尾部,同时挤掉等长的最旧部分——cache 总长度恒为 C,像一个固定大小的滑动窗口。
cache 像一个传送带:新内容从右边进来,等量的旧内容从左边出去。任意时刻 cache 里始终是最新的 C 个位置的 K/V。
当 C >= L(cache 容量能装下整条 prompt)时,窗口不会溢出,所有 chunk 的 KV 都保留,等价于普通 prefill。只有当 L > C(prompt 超过 cache 容量)时,滚动才会真正丢弃旧 KV——此时 chunk prefill 退化成一种"滑动窗口注意力",最旧的上下文会被遗忘,这是显存换精度的固有代价。
4. mask 切分:每个 chunk 的注意力范围
普通 prefill 的 mask 是一个L 的矩阵:query(查询)是全部 L 个位置,key(键)是整个 cache。chunk prefill 把 query 切成了 chunk,但 key 还是整个 cache(因为要 attend(关注)所有历史)。所以每个 chunk 的 mask 要:
从完整 mask 里切出这个 chunk 对应的 c 行(query 维);
- 右对齐到 cache 长度 C——因为 cache 是右对齐的(有效 KV 在尾部),mask 的 key 维也要右对齐,左边补 min_value。
这里的"右对齐":cache 有效内容在右端,mask 的 key 维就右对齐。
5. 把对齐和 chunk 串起来:完整 prefill 流程
几个细节点:
- padding_data 是左补左截函数,这里对齐目标不是固定的 L,而是 chunk_size 的整数倍(因为 chunk 必须等长,最后一块不能是零头),再 cap 到 max_kvcache_len。
aligned_len 向上取整到 chunk 倍数
position_ids 也要切:每个 chunk 的 position 是全序列 position 的一段,不能从 0 重新开始。
- 最后一个 chunk 给出第一个生成 token 的 logits:prefill 跑完所有 chunk 后,next_logits 就是下一个要生成的 token 的预测,从这里接 decode。
6. 对齐 vs Chunk Prefill:一张表厘清边界
这两个机制经常被混为一谈,其实正交:
对齐(padding/truncation) | Chunk Prefill | |
|---|---|---|
解决的问题 | 形状不匹配:导出期定形 vs 运行期变长 | 显存峰值:一次性算不动长序列 |
何时发生 | 喂进模型前,host 侧预处理 | prefill 计算阶段,循环内 |
对长度做什么 | 短了左补 pad、长了左截断 | 把已对齐序列切成多个 chunk 分批算 |
是否有损 | padding 无损、truncation 有损 | cache 够时无损、cache 溢出时丢旧 KV 有损 |
关键参数 | max_lm_input_len (L) | chunk_size (c)、max_kvcache_len (C) |
是否可独立开关 | 否(导出期固定,必须做) | 是(运行期策略,可关掉退回单次 prefill) |
产出 | (1, L) 定形输入 | 滚动 KV cache + 第一个生成 token |
7. 工程落地时tradeoff
落地时这几个值要一起调,互相牵制:
- max_lm_input_len (L):定形长度。设小了频繁触发截断(有损);设大了编译产物大、单 chunk 显存高。原则:覆盖业务最长 prompt,留少量余量。
- chunk_size (c):分批粒度。设小了循环次数多、吞吐低;设大了单次峰值显存高,可能仍 OOM。原则:让单 chunk 的 attention 激活 c x C 落在板端预算内,再尽量往大取。
- max_kvcache_len (C):cache 容量。必须 C >= L;C - L 是 decode 可生成空间。原则:C = L + max_new_tokens,既不浪费也不溢出。
- L == C 的陷阱:decode 空间为 0,能读不能写。校准/编译时务必给 decode 留 gap。
- 截断优先级:宁可调大 L 让截断不触发,也不要默认依赖截断——尤其 VLM,图像 token 在前部,左截断会砍掉图像信息。
总结

- 切:把对齐后的 L 个 token 按 chunk_size (c) 切成多块;
- 滚:每块前向后,用 torch.cat([old, new])[:, refresh_len:] 把新 K/V 滚动写进固定容量 C 的 KV cache,单次 attention 峰值从 L x L 降到 c x C;
- 对齐 mask:每个 chunk 的 causal mask(因果掩码)沿 query 维切分后,key 维右对齐到 C(左补 min_value),与 cache 右对齐一致;
- 串接对齐:先用 padding_data 对齐到 chunk 倍数,再进入 chunk 循环,最后一块给出第一个生成 token。
