prefix_grouper_utils.py — 该文件实现了 前缀共享优化(Prefix Sharing) 的工具函数¶
文件概述¶
模块路径: verl.trainer.ppo.prefix_grouper_utils
该文件实现了 前缀共享优化(Prefix Sharing) 的工具函数。当同一个 prompt 有多个不同的回复时(如 GRPO 场景),所有回复共享相同的 prompt 前缀。通过 PrefixGrouper,可以只计算一次前缀的 attention,然后分别处理不同的回复后缀,从而显著减少计算量。
在训练流程中的位置¶
这是一个性能优化模块,在 Actor 进行前向传播(计算 log_prob、entropy)时使用。当配置 actor.use_prefix_grouper=True 时启用。
关键代码讲解¶
1. 构建 Position IDs¶
def build_position_ids_for_prefix_grouper(prefix_grouper: PrefixGrouper) -> torch.Tensor:
"""为 PrefixGrouper 构建 position_ids,每个回复从 prefix_len 处重新开始"""
position_ids = torch.zeros(num_samples, max_len, dtype=torch.long, device=device)
for i, group in enumerate(prefix_grouper.group_info):
prefix_len = group.prefix_len
# 前缀部分:0, 1, 2, ..., prefix_len-1
position_ids[i, :prefix_len] = torch.arange(prefix_len, device=device)
cur_pos = prefix_len
for suffix_len in group.suffix_lens:
if suffix_len > 0:
# 每个后缀从 prefix_len 开始:prefix_len, prefix_len+1, ...
position_ids[i, cur_pos : cur_pos + suffix_len] = torch.arange(
prefix_len, prefix_len + suffix_len, device=device
)
cur_pos += suffix_len
return position_ids
这里的关键是:每个回复后缀的 position_ids 都从 prefix_len 开始,而不是从 prefix_len + 上一个后缀长度 开始。这样每个后缀独立地"续写"前缀。
2. 从 micro_batch 构建 PrefixGrouper¶
def build_pg_from_micro_batch(micro_batch, pad_token_id, padding_mode="right"):
"""从包含 prompts, responses, response_mask, uid 的 micro_batch 构建 PrefixGrouper"""
# 1. 根据 uid 确定分组大小
uids = micro_batch["uid"]
group_sizes = []
cur = 1
for i in range(1, bs):
if uids[i] == uids[i - 1]:
cur += 1
else:
group_sizes.append(cur)
cur = 1
group_sizes.append(cur)
# 2. 提取每组的前缀(第一个样本的 prompt)
prefix_indices = []
cursor = 0
for gs in group_sizes:
prefix_indices.append(cursor)
cursor += gs
prefix_ids = prompts.index_select(0, torch.tensor(prefix_indices))
# 3. 创建 PrefixGrouper
prefix_grouper = PrefixGrouper.from_ungrouped_masks(
prefix_mask=prefix_mask,
suffix_mask=response_mask,
group_sizes=group_sizes,
)
# 4. 拼接前缀和后缀的 input_ids
concat_input_ids = prefix_grouper.concat_input(prefix_ids, prefix_mask, responses, response_mask)
attention_mask = prefix_grouper.padding_mask
position_ids = build_position_ids_for_prefix_grouper(prefix_grouper)
return prefix_grouper, concat_input_ids, attention_mask, position_ids, responses, response_mask
3. 前向传播¶
def pg_forward(model, prefix_grouper, concat_input_ids, attention_mask, position_ids,
completion_ids, completion_mask, *, temperature=1.0, ...):
"""使用 PrefixGrouper 进行前向传播"""
# 1. 模型前向传播(内部使用 prefix_grouper 优化 attention 计算)
logits = model(
input_ids=concat_input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
use_cache=False,
prefix_grouper=prefix_grouper,
).logits
# 2. 将输出分离为前缀和后缀部分
prefix_out, prefix_mask, suffix_out_raw, suffix_mask_raw = prefix_grouper.split_output(
logits, include_prefix_last=1
)
# 3. 计算 log_probs
suffix_out = suffix_out_raw[:, :-1].float()
suffix_out /= temperature
log_probs = logprobs_from_logits(suffix_out, completion_ids_right)
# 4. 计算 entropy(如果需要)
entropy = None
if calculate_entropy and entropy_fn is not None:
entropy = entropy_fn(suffix_out)
return log_probs, entropy, suffix_mask
4. 完整的 micro_batch 前向传播¶
def forward_micro_batch_with_prefix_grouper(micro_batch, model, temperature, ...):
"""完整的前缀共享前向传播流程"""
# 1. 构建 PrefixGrouper
(prefix_grouper, concat_input_ids, attention_mask, position_ids,
responses, response_mask) = build_pg_from_micro_batch(micro_batch, pad_token_id)
# 2. 前向传播
with torch.autocast(device_type=device_name, dtype=param_dtype):
log_probs, entropy, suffix_mask = pg_forward(
model=model, prefix_grouper=prefix_grouper, ...
)
# 3. 零填充 padding 位置
padding_mask = suffix_mask == 0
log_probs = log_probs.masked_fill(padding_mask, 0.0)
# 4. 如果需要,补齐到目标长度
if log_probs.size(1) != target_response_length:
full_log_probs = log_probs.new_zeros(batch_size, target_response_length)
full_log_probs[:, :current_len] = log_probs
log_probs = full_log_probs
return entropy, log_probs
前缀共享示意图¶
不使用前缀共享:
Prompt A + Response 1: [Attention 计算完整序列]
Prompt A + Response 2: [Attention 计算完整序列] ← 前缀部分重复计算!
Prompt A + Response 3: [Attention 计算完整序列]
使用前缀共享:
Prompt A (前缀): [Attention 计算一次]
+-- Response 1 (后缀): [续写 Attention]
+-- Response 2 (后缀): [续写 Attention] ← 前缀只计算一次!
+-- Response 3 (后缀): [续写 Attention]
核心类/函数列表¶
| 名称 | 类型 | 作用 |
|---|---|---|
build_position_ids_for_prefix_grouper() |
function | 构建考虑前缀共享的 position_ids |
build_pg_from_micro_batch() |
function | 从 micro_batch 创建 PrefixGrouper |
pg_forward() |
function | 使用 PrefixGrouper 进行模型前向传播 |
forward_micro_batch_with_prefix_grouper() |
function | 完整的前缀共享前向传播封装 |
数据流和调用关系¶
Actor Worker (compute_log_prob 或 update_actor)
|
+-- forward_micro_batch_with_prefix_grouper()
|
+-- build_pg_from_micro_batch()
| +-- 按 uid 分组
| +-- PrefixGrouper.from_ungrouped_masks()
| +-- prefix_grouper.concat_input()
| +-- build_position_ids_for_prefix_grouper()
|
+-- pg_forward()
+-- model(prefix_grouper=prefix_grouper)
+-- prefix_grouper.split_output()
+-- logprobs_from_logits()
小结¶
prefix_grouper_utils.py 是一个重要的性能优化模块:
- 计算节省:同一 prompt 的多个回复共享前缀的 attention 计算
- 内存节省:减少重复的 KV cache 存储
- 透明集成:通过
PrefixGrouper库提供的 API,对模型架构的侵入性最小 - 自动分组:根据
uid自动识别同一 prompt 的回复
在 GRPO 等需要同一 prompt 多次采样的场景中,这个优化可以显著提升训练速度。