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 的规则化聚合机制确保了在多步训练场景下指标的正确计算。