跳转至

detach_utils.py — 提供全异步训练的数据处理工具

文件路径: verl/experimental/fully_async_policy/detach_utils.py

文件概述

提供全异步训练的数据处理工具,包括: 1. RolloutSample / ValidateMetrics - 数据结构定义 2. prepare_single_generation_data - 单样本数据预处理 3. assemble_batch_from_rollout_samples - 将多个 rollout 样本组装为训练 batch 4. MetricsAggregator - 跨多个训练步聚合指标

关键代码讲解

1. RolloutSample 数据结构

@dataclass
class RolloutSample:
    full_batch: Any                          # 原始 batch 信息
    agent_loop_output_list: list[AgentLoopOutput]  # 生成的输出
    sample_id: str                           # 样本 ID
    epoch: int                               # 当前 epoch
    processing_times: list[float]            # 处理时间
    tool_calls: list[float]                  # 工具调用时间
    param_version: int                       # 使用的参数版本
    param_version_start: list[int]           # 各轮开始时的参数版本
    param_version_end: list[int]             # 各轮结束时的参数版本
    rollout_status: dict[str, Any]           # Rollouter 状态统计

2. 单样本数据准备

def prepare_single_generation_data(batch_dict, config) -> DataProto:
    full_batch = DataProto.from_single_dict(batch_dict)

    # 根据配置选择 Agent 类型
    if config.actor_rollout_ref.rollout.multi_turn.enable:
        full_batch.non_tensor_batch["agent_name"] = np.array(
            ["async_partial_tool_agent"] * len(full_batch), dtype=object
        )
    else:
        full_batch.non_tensor_batch["agent_name"] = np.array(
            ["partial_single_turn_agent"] * len(full_batch), dtype=object
        )

    # 重复 n 次(n 个 rollout 采样)
    full_batch = full_batch.repeat(repeat_times=config.actor_rollout_ref.rollout.n, interleave=True)
    return full_batch

3. 批量组装

def assemble_batch_from_rollout_samples(rollout_samples, tokenizer, config, balance_batch=None):
    # 合并所有样本
    rollout_samples_batch = [rs.full_batch for rs in rollout_samples]
    final_batch = DataProto.concat(rollout_samples_batch)

    # 计算 response_mask
    if "response_mask" not in final_batch.batch.keys():
        final_batch.batch["response_mask"] = compute_response_mask(final_batch)

    # 收集统计信息
    processing_time_stats = {
        "processing_time/avg": np.mean(processing_times),
        "processing_time/max": np.max(processing_times),
        "processing_time/tp99": np.percentile(processing_times, 99),
        ...
    }

    # 统计部分回滚信息
    param_version_diff = [abs(a - b) for a, b in zip(param_version_end, param_version_start)]
    partial_stats = {
        "fully_async/partial/total_partial_num": ...,
        "fully_async/partial/partial_ratio": ...,
    }

    return final_batch

4. MetricsAggregator - 指标聚合器

在全异步模式下,一个参数版本可能对应多个训练步,需要将这些步的指标聚合后再记录:

class MetricsAggregator:
    def __init__(self, total_gpus):
        self.metric_values: dict[str, list[float]] = defaultdict(list)
        self.aggregation_rules = self._init_aggregation_rules()

    def add_step_metrics(self, metrics, sample_count, timestamp=None):
        """添加单步指标"""
        for key, value in metrics.items():
            self.metric_values[key].append(float(value))

    def get_aggregated_metrics(self):
        """获取聚合后的指标"""
        aggregated = {}
        for metric_name, values in self.metric_values.items():
            aggregated[metric_name] = self._aggregate_single_metric(metric_name, values)
        return aggregated

聚合规则示例: - 时间类指标用 sum - 吞吐量指标用 avg - 计数类指标用 last(取最后一个值) - 极值类指标用 max / min

核心类/函数列表

名称 类型 说明
RolloutSample 数据类 单个 rollout 样本的完整信息
ValidateMetrics 数据类 验证指标
prepare_single_generation_data() 函数 准备单样本的生成数据
assemble_batch_from_rollout_samples() 函数 组装训练 batch
MetricsAggregator 类 多步指标聚合器

与其他模块的关系

  • 被 FullyAsyncRollouter 使用(prepare_single_generation_data、RolloutSample)
  • 被 FullyAsyncTrainer 使用(assemble_batch_from_rollout_samples、MetricsAggregator)

小结

detach_utils.py 是全异步训练的数据处理层。它解决了异步环境下的数据格式转换(单样本 -> 训练 batch)、统计指标收集和聚合等问题。MetricsAggregator 的规则化聚合机制确保了在多步训练场景下指标的正确计算。