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 提供了全面的训练监控能力:
- 数据质量监控:奖励分数、序列长度、中止比例
- 性能监控:各阶段耗时、吞吐量
- 训练稳定性监控:Critic 解释方差、梯度方差代理
- 验证评估:pass@k、best@k、majority voting 等高级统计
这些指标对于调试训练问题和调优超参数至关重要。