跳转至

metric_utils.py — 该文件实现了 PPO 训练过程中的指标计算工具

文件概述

模块路径: verl.trainer.ppo.metric_utils

该文件实现了 PPO 训练过程中的指标计算工具。它提供了丰富的监控指标,包括奖励统计、优势统计、序列长度统计、时间性能指标、吞吐量指标和方差代理指标等。

在训练流程中的位置

在 ray_trainer.py 的 fit() 主循环中,每个训练步骤结束后都会调用这些函数来收集指标,然后发送到日志系统(如 WandB)。

关键代码讲解

1. 数据指标计算

def compute_data_metrics(batch: DataProto, use_critic: bool = True) -> dict[str, Any]:
    """计算一个 batch 的各种统计指标"""

    # 序列级分数:对 token_level_scores 求和
    sequence_score = batch.batch["token_level_scores"].sum(-1)
    sequence_reward = batch.batch["token_level_rewards"].sum(-1)

    # 检测被中止的样本(response_length == 0)
    aborted_mask = (response_length == 0).bool()
    non_aborted_mask = ~aborted_mask

    # 对非中止样本计算分数统计
    non_aborted_sequence_score = sequence_score[non_aborted_mask]
    score_mean = torch.mean(non_aborted_sequence_score).detach().item()
    score_max = torch.max(non_aborted_sequence_score).detach().item()
    score_min = torch.min(non_aborted_sequence_score).detach().item()

    # 优势和回报统计(只在有效 token 上计算)
    valid_adv = torch.masked_select(advantages, response_mask)
    valid_returns = torch.masked_select(returns, response_mask)

    # Critic 相关指标
    if use_critic:
        valid_values = torch.masked_select(values, response_mask)
        # 价值函数解释方差:越接近 1 说明 Critic 越好
        return_diff_var = torch.var(valid_returns - valid_values)
        return_var = torch.var(valid_returns)
        vf_explained_var = 1.0 - return_diff_var / (return_var + 1e-5)

返回的指标字典包括: - critic/score/mean|max|min - 奖励分数统计 - critic/advantages/mean|max|min - 优势统计 - critic/vf_explained_var - Critic 的解释方差(衡量 Critic 质量) - response_length/mean|max|min|clip_ratio - 回复长度统计 - prompt_length/mean|max|min - 提示长度统计 - response/aborted_ratio - 被中止样本的比例

2. 时间指标计算

def compute_timing_metrics(batch: DataProto, timing_raw: dict[str, float]) -> dict[str, Any]:
    """计算各阶段的时间和每 token 耗时"""
    num_prompt_tokens = torch.sum(response_info["prompt_length"]).item()
    num_response_tokens = torch.sum(response_info["response_length"]).item()
    num_overall_tokens = num_prompt_tokens + num_response_tokens

    # 不同阶段使用不同的 token 数做归一化
    num_tokens_of_section = {
        "gen": num_response_tokens,                    # 生成只计回复 token
        "ref": num_overall_tokens,                     # 其他阶段计全部 token
        "values": num_overall_tokens,
        "adv": num_overall_tokens,
        "update_critic": num_overall_tokens,
        "update_actor": num_overall_tokens,
    }

    return {
        **{f"timing_s/{name}": value for name, value in timing_raw.items()},
        **{f"timing_per_token_ms/{name}": timing_raw[name] * 1000 / num_tokens
           for name in set(num_tokens_of_section.keys()) & set(timing_raw.keys())},
    }

3. 吞吐量指标

def compute_throughout_metrics(batch, timing_raw, n_gpus):
    """计算每 GPU 每秒处理的 token 数(吞吐量)"""
    total_num_tokens = sum(batch.meta_info["global_token_num"])
    time = timing_raw["step"]
    return {
        "perf/total_num_tokens": total_num_tokens,
        "perf/time_per_step": time,
        "perf/throughput": total_num_tokens / (time * n_gpus),  # tokens/s/GPU
    }

4. 方差代理指标

def compute_variance_proxy_metrics(batch, gradient_norm=None):
    """
    计算梯度方差的代理指标(不需要实际计算所有样本的梯度)

    理论基础:
    - Proxy 1 (Signal): ||g_mean||^2 = gradient_norm^2
    - Proxy 2 (Total Power): E[A^2 * W] where W 是 score norm 代理
    - Proxy 3 (Noise): (Proxy2 - Proxy1) / (N-1)
    """
    # W(τ) = Σ_t[1 - 2π_t(y_t) + Σπ²]
    pi_t = torch.exp(batch.batch["old_log_probs"])
    w_per_timestep = 1 - 2 * pi_t + batch.batch["sum_pi_squared"]

    # 如果有 IS 权重,用 ρ² 缩放 W
    if "rollout_is_weights" in batch.batch:
        rollout_is_weights = batch.batch["rollout_is_weights"]
        w_per_timestep = w_per_timestep * (rollout_is_weights**2).detach()

    # Proxy 2: E[A^2 * W]
    proxy2_total_power = (advantages_scalar**2 * w_values_clamped).mean()

    # Proxy 3: (Proxy2 - Proxy1) / (N-1)
    if proxy1_signal_strength is not None and batch_size > 1:
        proxy3_pure_noise = (1.0 / (batch_size - 1)) * (proxy2_total_power - proxy1_signal_strength)

方差代理指标非常高级,它帮助我们在不增加计算成本的情况下监控梯度方差,这对于理解训练稳定性很有价值。

5. 验证指标处理

def process_validation_metrics(data_sources, sample_uids, infos_dict, seed=42):
    """
    处理验证集指标,计算 pass@k, best@k, majority@k 等统计
    """
    # 按数据源和 prompt 分组
    for sample_idx, data_source in enumerate(data_sources):
        uid = sample_uids[sample_idx]
        for var_name, var_vals in infos_dict.items():
            data_src2uid2var2vals[data_source][uid][var_name].append(var_vals[sample_idx])

    # 对每组计算 bootstrap 统计
    for n in ns:
        # 计算 best@n (最好的 n 个中取最大)
        (bon_mean, bon_std), (won_mean, won_std) = bootstrap_metric(
            data=var_vals, subset_size=n,
            reduce_fns=[np.max, np.min], ...
        )
        # 计算 majority@n (多数投票)
        if has_pred:
            [(maj_n_mean, maj_n_std)] = bootstrap_metric(
                data=vote_data, subset_size=n,
                reduce_fns=[partial(calc_maj_val, vote_key="pred", val_key="val")],
            )

6. Bootstrap 采样

def bootstrap_metric(data, subset_size, reduce_fns, n_bootstrap=1000, seed=42):
    """
    通过 bootstrap 重采样估计指标的均值和标准差

    1. 从数据中随机有放回抽取 subset_size 个样本,重复 n_bootstrap 次
    2. 对每次抽样应用 reduce_fns(如 np.max, np.min)
    3. 计算所有次数的均值和标准差
    """
    bootstrap_idxs = np.random.choice(n_data, size=(n_bootstrap, subset_size), replace=True)

    for fn_idx, reduce_fn in enumerate(reduce_fns):
        for boot_idx in range(n_bootstrap):
            sample = data_np[bootstrap_idxs[boot_idx]]
            metric_results[fn_idx, boot_idx] = reduce_fn(sample)

    return [(np.mean(results), np.std(results)) for results in metric_results]

核心类/函数列表

名称 类型 作用
compute_data_metrics() function 计算奖励/优势/长度等数据指标
compute_timing_metrics() function 计算各阶段时间和每 token 耗时
compute_throughout_metrics() function 计算吞吐量 (tokens/s/GPU)
compute_variance_proxy_metrics() function 计算梯度方差代理指标
process_validation_metrics() function 处理验证集指标(pass@k, best@k 等)
bootstrap_metric() function Bootstrap 重采样估计统计量
calc_maj_val() function 多数投票计算
reduce_metrics() function 将指标列表取均值(已废弃)

数据流和调用关系

ray_trainer.py: fit()
    |
    +-- compute_data_metrics(batch)
    |       |-- 计算 critic/score/*, critic/rewards/*, critic/advantages/*
    |       |-- 计算 response_length/*, prompt_length/*
    |       +-- 计算 critic/vf_explained_var (如果用 Critic)
    |
    +-- compute_timing_metrics(batch, timing_raw)
    |       +-- 计算 timing_s/*, timing_per_token_ms/*
    |
    +-- compute_throughout_metrics(batch, timing_raw, n_gpus)
    |       +-- 计算 perf/throughput, perf/time_per_step
    |
    +-- compute_variance_proxy_metrics(batch, gradient_norm)
    |       +-- 计算 variance_proxy/proxy1_signal_strength 等
    |
    +-- _validate() --> process_validation_metrics()
            +-- 计算 val-core/*, val-aux/* (best@k, maj@k 等)

小结

metric_utils.py 提供了全面的训练监控能力:

  1. 数据质量监控:奖励分数、序列长度、中止比例
  2. 性能监控:各阶段耗时、吞吐量
  3. 训练稳定性监控:Critic 解释方差、梯度方差代理
  4. 验证评估:pass@k、best@k、majority voting 等高级统计

这些指标对于调试训练问题和调优超参数至关重要。