Continuous Batching 的工作原理
LLM 中的批量推理
深度学习模型进行推理时,通常会同时处理一个 batch 的输入,以最大化硬件利用率和吞吐量。以分类任务为例,输入可能是一个向量,但如果每次只处理一个输入,计算就会变成多次向量-矩阵乘法。对于大模型而言,一次向量-矩阵乘法需要将整个矩阵加载到内存中,效率可能较低。若模型同时处理一个 batch 的输入,则可以进行矩阵-矩阵乘法,通常效率更高。
在 LLM 中,输入通常是一个由多个序列组成的 batch,而每个序列的长度各不相同。传统批处理方法要求 batch 中的所有输入具有相同长度。通常的做法是为较短的序列补充特殊 token,使它们与 batch 中最长的序列对齐。
计算注意力分数时,模型使用 mask 忽略 padding token,确保它们不会影响注意力机制。
不过,使用 padding 会浪费计算资源。假设 batch 中一个序列的长度为 10,另一个序列的长度为 100,那么在计算注意力分数时,模型需要计算一个 的注意力矩阵。其中,只有 个元素对应有效 token,其余计算都与 padding 有关,因而没有实际意义。
此外,事先无法知道 batch 中每个序列会生成多少个 token,因此序列长度会在推理过程中不断增加。为了避免频繁拷贝 KV Cache,通常需要为每个序列预先分配固定大小的缓冲区。这同样会浪费内存:有些序列可能很快结束生成,但其 KV Cache 仍占用了大量空间。
什么是 Continuous Batching
Continuous batching 允许模型在一个 batch 中处理长度各异的序列,并避免使用大量 padding。其基本思路是将多个序列拼接为一个长序列,同时记录每个原始序列在 batch 中的起止位置。
下图是 continuous batching 的简化示意图:
上图中,我用不同颜色表示 batch 中的不同序列。使用 continuous batching 后,实际上只有一个序列,只是这个序列由多个原始序列拼接而成。
基于 Transformer 的 LLM 由 Embedding 层、多头自注意力层、前馈层等组成。Embedding 层和前馈层对每个 token 独立计算,因此天然可以处理连续拼接后的 batch。对于多头自注意力层,则需要额外保证不同序列的 token 不会相互影响。由于已经记录了每个序列的起止位置,只需将注意力计算限制在各自序列的边界内即可。
拼接后的序列经过模型后,可以根据每个序列的结束位置提取对应的输出。
每次推理后,系统都会判断是否有序列已生成完毕(例如生成了序列结束 token 或达到最大长度),并将已完成的序列从下一步待处理的 batch 中移除。
这就是 LLM 中 continuous batching 的核心思想:高效地将不同长度的序列拼接为一个 batch,最大化硬件利用率,同时避免不同序列之间相互影响。
如何实现 Continuous Batching
要实现 continuous batching,需要仔细管理序列及其位置。可以维护当前活跃序列的长度列表;进行注意力计算时,再将总序列按原始序列切分为多个片段,并按传统方式分别计算注意力。
阅读以下代码片段时,可参考这个 notebook 中的完整实现。
使用 Continuous Batching 计算注意力
可以维护一个累积长度列表,记录每个序列在拼接 batch 中的起止位置。
cu_lens = [0]
for seq in sequences:
cu_lens.append(cu_lens[-1] + len(seq))
例如,输入 batch 中有三个序列,长度分别为 3、5、2,则 cu_lens 为 [0, 3, 8, 10]。
计算注意力时,可以遍历每个序列片段,分别在子序列内部计算注意力,或者使用如下图所示的方式,使用 mask 将不同序列之间的注意力分数置为负无穷。
在生产环境中,通常会使用更高效的实现方式,例如 Flash Attention。Flash Attention 使用 cu_seqlens 张量记录每个序列的累积长度,并据此将总序列切分为多个片段;每个子序列由一个线程块(thread block)计算注意力。
KV Cache 的管理
在推理过程中,通常会使用 KV Cache 来缓存每个序列的 key 和 value,以便在生成下一个 token 时复用。对于 continuous batching,需要为每个序列维护独立的 KV Cache,并在每次生成新 token 时,将其追加到对应序列的 KV Cache 中。
实际实现通常采用分页方式管理 KV Cache。每个序列的 KV Cache 可以存储在不连续的内存页中,并通过索引访问。
Continuous Batching 中的旋转位置编码
应用旋转位置编码时,也需要考虑每个 token 在原始序列中的位置。可以维护一个 positions 张量,记录拼接后的 batch 中每个 token 在其原始序列中的位置。
positions = []
for seq in sequences:
seq_len = len(seq)
positions.extend(range(seq_len))
例如,输入 batch 中有三个序列,长度分别为 3、5、2,则 positions 为 [0, 1, 2, 0, 1, 2, 3, 4, 0, 1]。
q = self.rotary_embedding(q, positions)
k = self.rotary_embedding(k, positions)
可以将 positions 张量作为索引,为拼接 batch 中的每个 token 计算正确的位置编码。
cos = self.cos_cached[positions]
sin = self.sin_cached[positions]
生成下一个 Token
在输出层中,可以根据各序列的结束位置提取其最后一个 token 对应的 logits。
class Qwen3ForCausalLM(nn.Module):
def __init__(self, config: Qwen3Config):
super().__init__()
self.model = Qwen3Model(config)
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
def forward(self, input_ids: torch.Tensor, positions: torch.Tensor, cu_lens: torch.Tensor) -> torch.Tensor:
# [seqlen, hidden_size]
x = self.model(input_ids, positions, cu_lens)
# [seqlen, vocab_size]
# 提取每个序列最后一个 token 的输出
x = self.lm_head(x[cu_lens[1:]-1, :])
return x
Batch 管理
通过 continuous batching,可以高效地对多个不同长度的序列执行推理。系统维护一个活跃序列列表,为每个序列生成下一个 token 后,将已完成生成的序列从列表中移除。
下面是管理 batch 过程的简化示例。
首先定义 Request 类,用于保存每个序列的 token。
class Request:
"""
一个请求包含待处理序列的 token。
"""
def __init__(self, tokens):
self.tokens = tokens
每次获得一批请求后,拼接它们的 token,并维护相应的位置和累积长度,然后执行一步生成。
def generate_one_step(model, requests: list[Request]):
"""
为 batch 中每个请求生成一个 token。
"""
tokens = []
positions = []
cu_lens = [0]
for req in requests:
tokens.extend(req.tokens)
positions.extend(range(len(req.tokens)))
cu_lens.append(cu_lens[-1] + len(req.tokens))
tokens = torch.tensor(tokens, dtype=torch.long, device=device)
positions = torch.tensor(positions, dtype=torch.long, device=device)
cu_lens = torch.tensor(cu_lens, dtype=torch.long, device=device)
# [len(requests), vocab_size]
logits = model(tokens, positions, cu_lens)
next_tokens = torch.argmax(logits, dim=-1)
return next_tokens.tolist()
在主生成循环中,系统持续为这批请求生成 token,并将新 token 追加到对应请求中。如果某个请求生成了序列结束 token,就将其从 batch 中移除。
def generate(model: Qwen3ForCausalLM, tokenizer, prompts: list[str], enable_think=True, max_new_tokens=64):
"""
使用 continuous batching 为一批 prompt 生成 token。
如果一个请求已经完成,则将其从 batch 中移除。
"""
requests = []
for prompt in prompts:
prompt = apply_chat_template(prompt, enable_think)
tokens = qwen3_tokenizer.encode(prompt).ids
req = Request(tokens)
requests.append(req)
eos_token = tokenizer.encode("<|im_end|>").ids[0]
new_tokens = 0;
while len(requests) and new_tokens < max_new_tokens:
new_tokens += 1
tokens = generate_one_step(model, requests)
for req, token in zip(requests, tokens):
req.tokens.append(token)
# 移除已经完成的请求
requests = [req for req in requests if req.tokens[-1] != eos_token]
return requests
总结
本文介绍了 continuous batching 在大语言模型(LLM)中的工作原理。我在这个 notebook 中实现了一个简化版本的 continuous batching。希望这些说明和代码片段能够帮助你理解 LLM 中 continuous batching 的概念。