跳转至

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 是一个重要的性能优化模块:

  1. 计算节省:同一 prompt 的多个回复共享前缀的 attention 计算
  2. 内存节省:减少重复的 KV cache 存储
  3. 透明集成:通过 PrefixGrouper 库提供的 API,对模型架构的侵入性最小
  4. 自动分组:根据 uid 自动识别同一 prompt 的回复

在 GRPO 等需要同一 prompt 多次采样的场景中,这个优化可以显著提升训练速度。