专栏算法工具链【大模型】Chunk Prefill——把一次算不动的事拆成多次

【大模型】Chunk Prefill——把一次算不动的事拆成多次

no_name2026-07-26
32
0

声明:本文主要参考开源资料进行学习整理,如有错漏,欢迎评论交流~

引言

假设已经解决了"形状不匹配"的问题——把任意长度的 prompt 左 padding / 左截断对齐到固定长度 L,喂进定形的 prefill 静态图。但形状对齐了,显存不一定够。设想一条对齐到 L=1216 的 prompt,prefill 阶段要算一个 1216 x 1216 的注意力矩阵;如果业务允许的 L 开到几千,单次 prefill 的激活显存峰值会直接顶爆板端 DDR/L2,OOM,这就是显存峰值墙
chunk prefill登场,把一次算不动的长序列,切成 chunk 分批喂进去,靠滚动 KV cache 化解峰值。

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。

Chunk prefill 的思路:与其一次性喂 L 个 token,不如把 L 个 token 切成若干个 chunk,每个 chunk 大小 c,分多次喂进 prefill 模型,每次只算 c 个 token 的 attention,滚动更新 KV cache。 这样峰值显存大幅下降,长 prompt 也能跑。

代价是吞吐:原来 1 次前向完成,现在要 L/c 次。但换来了"能跑"——这是显存受限的端侧场景下的必要取舍。

2. 前置:KV cache 的容量与右对齐

chunk prefill 全程围绕 KV cache 转,先把它的容量和"右对齐"约定讲清楚。

KV cache 是一个长度恒为 C(max_kvcache_len)的固定缓冲,存放每一层注意力的 K 和 V。它的两个关键性质:
  1. 容量固定为 C:prefill 的 L 个位置 + decode 生成的新位置,总占用不得超过 C。 L ,且 C - L 是 decode 可生成空间。
  2. 有效内容右对齐:真实 token 的 K/V 总是写在 cache 的右端(高位)。这源自于“左 padding”——真实 token 在序列右端,进 cache 后自然落在右端。后续 decode 新 token 也是追加在右端往后递增。例如KV cache (长度 C=8,当前有 5 个有效 token):

有这两点先验知识,下面的滚动写入和 mask 切分才有据可依。

3. 核心机制:滚动 KV cache

chunk prefill 的灵魂在滚动 KV cache
每个 chunk 算出的 K/V,追加到 KV cache 的尾部,同时挤掉等长的最旧部分——cache 总长度恒为 C,像一个固定大小的滑动窗口。
逐行解释这一行 torch.cat([old, new], dim=1)[:, refresh_len:]:

cache 像一个传送带:新内容从右边进来,等量的旧内容从左边出去。任意时刻 cache 里始终是最新的 C 个位置的 K/V。

当 C >= L(cache 容量能装下整条 prompt)时,窗口不会溢出,所有 chunk 的 KV 都保留,等价于普通 prefill。只有当 L > C(prompt 超过 cache 容量)时,滚动才会真正丢弃旧 KV——此时 chunk prefill 退化成一种"滑动窗口注意力",最旧的上下文会被遗忘,这是显存换精度的固有代价。

4. mask 切分:每个 chunk 的注意力范围

chunk prefill 还有一个绕不开的难点:causal mask(因果掩码)怎么切

普通 prefill 的 mask 是一个L 的矩阵:query(查询)是全部 L 个位置,key(键)是整个 cache。chunk prefill 把 query 切成了 chunk,但 key 还是整个 cache(因为要 attend(关注)所有历史)。所以每个 chunk 的 mask 要:

  1. 从完整 mask 里切出这个 chunk 对应的 c 行(query 维);

  2. 右对齐到 cache 长度 C——因为 cache 是右对齐的(有效 KV 在尾部),mask 的 key 维也要右对齐,左边补 min_value。
为什么是右对齐?因为 KV cache 的有效内容在尾部(滚动写入时新 KV 追加在右端)。mask 的 key 维度必须和 cache 对齐,否则 chunk 的 query 会 attend(关注)到错误的位置。每个 chunk mask 的形状是 (bs, 1, c, C),左补 min_value 把有效部分推到右边。

这里的"右对齐":cache 有效内容在右端,mask 的 key 维就右对齐。

5. 把对齐和 chunk 串起来:完整 prefill 流程

对齐和 chunk prefill 不是二选一,而是串联的:先用 padding_data 把序列对齐,再切成 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

一句话:对齐是"形状归一化",chunk prefill 是"计算分批化"。先对齐,再分批。

7. 工程落地时tradeoff

落地时这几个值要一起调,互相牵制:

  1. max_lm_input_len (L):定形长度。设小了频繁触发截断(有损);设大了编译产物大、单 chunk 显存高。原则:覆盖业务最长 prompt,留少量余量。
  2. chunk_size (c):分批粒度。设小了循环次数多、吞吐低;设大了单次峰值显存高,可能仍 OOM。原则:让单 chunk 的 attention 激活 c x C 落在板端预算内,再尽量往大取。
  3. max_kvcache_len (C):cache 容量。必须 C >= L;C - L 是 decode 可生成空间。原则:C = L + max_new_tokens,既不浪费也不溢出。
  4. L == C 的陷阱:decode 空间为 0,能读不能写。校准/编译时务必给 decode 留 gap。
  5. 截断优先级:宁可调大 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。
max_lm_input_len / chunk_size / max_kvcache_len 三个参数的牵制关系理清,prefill 的显存墙就已翻过。

后续

chunk prefill 把长序列拆成 chunk 分批算,但每个 chunk 喂进 attention 时,那张 mask 是怎么构造的——attention_mask 一维的 0/1 标记怎么变成 L x C 的屏蔽矩阵、causal(因果,不能看未来)和 pad 屏蔽如何叠加、为什么屏蔽位填 -32768 而非 -inf——这些会在 causal attention(因果注意力)与 causal_mask(因果掩码)一篇里展开。
算法工具链
社区征文杂谈前沿技术
评论0
0/600