损失函数
PyTRIO 为监督学习和强化学习提供了内置的损失函数。
你可以通过将字符串传递给 forward_backward 来选择损失函数:
future = training_client.forward_backward(
data,
loss_fn="cross_entropy",
)
result = future.result()内置损失函数
目前 PyTRIO 支持的内置损失函数如下:
| 损失函数 | 适用场景 | 说明 |
|---|---|---|
cross_entropy | 监督学习 | 标准交叉熵损失,适用于分类任务。以模型输出的 logits 和目标标签计算负对数似然。 |
importance_sampling | 离线强化学习 | 使用重要性采样对 off-policy 数据进行修正,通过行为策略与目标策略的概率比值对梯度加权。 |
ppo | 在线强化学习 | Proximal Policy Optimization 损失,通过裁剪概率比值限制策略更新幅度,提升训练稳定性。 |
cispo | 在线/离线强化学习 | Clipped Importance Sampling Policy Optimization,通过裁剪后的重要性比值为策略梯度加权,适合异步或 off-policy 场景。 |
dro | 离线强化学习 | Direct Reward Optimization,在奖励项之外加入二次惩罚,约束策略相对采样策略的更新幅度。 |
cross_entropy
在监督学习中,我们实现了标准的交叉熵损失(即负对数似然),该损失优化策略 以最大化 token 的对数概率:
其中 weights 为 0 或 1,通常由 renderer.build_supervised_example() 生成,该函数返回 (model_input, weights)(即用于指定需要训练的目标助手轮次)。
其实现方式为:
# Apply weights and compute elementwise loss
elementwise_loss = -target_logprobs * weights
# Apply sum reduction to get the total loss
loss = elementwise_loss.sum() # scalarcross_entropy损失需要Datum的loss_fn_inputs中传入target_tokens和weights两个字段:
target_tokens: array[(N,), int] | array[(N, K), int]:target token IDsweights: array[(N,), float] | array[(N, K), float]:token级的损失权重(通常来自渲染器)
输出:
logprobs: array[(N,), float] | array[(N, K), float]:请求的target token的对数概率
指标:
loss_sum:SDK 返回的聚合损失,是一个标量
importance_sampling
对于强化学习,我们实现了策略梯度目标的一个常见变体,适用于学习策略 与采样策略 存在差异的实际场景(例如由于非确定性导致的 off-policy 情况)。
问题在于,若两者存在差异,则目标:
由于 (采样器)并不严格等同于期望的 (学习器),会导致估计有偏。为修正此偏差,我们采用改进的"重要性采样"目标:
该目标可得到正确的期望奖励。公式中:
- (
target_logprobs)来自学习器,在forward_backward的前向阶段计算。 - (
sampling_logprobs)来自采样器,在采样时记录,用作修正项。
其实现方式为:
# Compute probability ratio
prob_ratio = torch.exp(target_logprobs - sampling_logprobs)
# Compute importance-weighted loss
loss = -(prob_ratio * advantages).sum()importance_sampling 损失需要 Datum 的 loss_fn_inputs 中传入以下字段:
target_tokens: array[(N,), int]:target token IDs(来自采样器 )logprobs: array[(N,), float]:token 的sampling_logprobsadvantages: array[(N,), float]:RL 的优势值(正值表示强化,负值表示抑制)
输出:
logprobs: array[(N,), float]:token 的target_logprobs
指标:
loss_sum:SDK 返回的聚合损失,是一个标量
ppo
PPO(Schulman et al., 2017)通过引入裁剪目标函数来解决标准策略梯度方法的问题,将策略更新限制在采样分布的邻域内,从而防止在同一 rollout 分布上进行多步梯度更新时出现过大的策略偏移。
该目标函数通过裁剪重要性比值 来防止策略更新幅度过大,其中 为学习器策略, 为采样策略。注意,PPO 的裁剪与损失计算均以 token 为单位独立进行。
PPO 裁剪目标为:
最终 PPO 损失结合了裁剪与未裁剪两个目标:
其中 和 为超参数(当前在 PyTRIO 中固定为 0.2)。
其实现方式为:
# Compute probability ratio
prob_ratio = torch.exp(target_logprobs - sampling_logprobs)
# Apply clipping
clipped_ratio = torch.clamp(prob_ratio, clip_low_threshold, clip_high_threshold)
# Compute both objectives
unclipped_objective = prob_ratio * advantages
clipped_objective = clipped_ratio * advantages
# Take minimum (most conservative)
ppo_objective = torch.min(unclipped_objective, clipped_objective)
# PPO loss is negative of objective
loss = -ppo_objective.sum()ppo 损失需要 Datum 的 loss_fn_inputs 中传入以下字段:
target_tokens: array[(N,), int]:target token IDs(来自采样器 )logprobs: array[(N,), float]:token 的sampling_logprobsadvantages: array[(N,), float]:RL 的优势值
输出:
logprobs: array[(N,), float]:token 的target_logprobs
指标:
loss_sum:SDK 返回的聚合损失,是一个标量
ps:还可通过 loss_fn_config 自定义裁剪阈值:
fwd_bwd_future = await training_client.forward_backward_async(
data=data,
loss_fn="ppo",
loss_fn_config={"clip_low_threshold": 0.9, "clip_high_threshold": 1.1}
)
fwd_bwd_result = await fwd_bwd_futurecispo
CISPO(Clipped Importance Sampling Policy Optimization)是一种策略梯度方法。它与 PPO 都会使用重要性比值 ,但差异在于:PPO 裁剪的是目标函数本身,而 CISPO 会先裁剪重要性比值,并将裁剪后的比值作为 target_logprobs 的梯度系数。
CISPO 目标为:
其中 表示 stop-gradient。裁剪后的比值会被 detach,因此它只作为系数影响梯度大小,不会通过比值本身继续反向传播。
其实现方式为:
# Compute probability ratio
prob_ratio = torch.exp(target_logprobs - sampling_logprobs)
# Apply clipping
clipped_ratio = torch.clamp(prob_ratio, clip_low_threshold, clip_high_threshold)
# Compute CISPO objective (detach the clipped ratio)
cispo_objective = clipped_ratio.detach() * target_logprobs * advantages
# CISPO loss is negative of objective
loss = -cispo_objective.sum()cispo 损失需要 Datum 的 loss_fn_inputs 中传入以下字段:
target_tokens: array[(N,), int]:target token IDs(来自采样器 )logprobs: array[(N,), float]:token 的sampling_logprobsadvantages: array[(N,), float]:RL 的优势值
输出:
logprobs: array[(N,), float]:token 的target_logprobs
指标:
loss_sum:SDK 返回的聚合损失,是一个标量
CISPO 默认使用单侧裁剪:不启用下界,只限制上界,即 clip_low_threshold=0.0、clip_high_threshold=4.0。你也可以通过 loss_fn_config 显式设置:
fwd_bwd_future = await training_client.forward_backward_async(
data=data,
loss_fn="cispo",
loss_fn_config={"clip_low_threshold": 0.0, "clip_high_threshold": 4.0}
)在异步或 off-policy 训练中,采样策略 可能滞后于当前学习策略 。如果设置较高的下界,会让已经被当前策略降低概率的旧 token 仍然获得较大的梯度系数,削弱重要性采样对 stale token 的衰减效果。因此,默认关闭下界通常更稳健。
dro
DRO(Direct Reward Optimization)是一种通用的 off-policy 甚至 offline 强化学习方法。它在奖励加权的 logprob 项之外加入二次惩罚,约束当前学习策略 不要相对采样策略 偏移过大。
DRO 目标为:
其中 控制二次惩罚强度。注意,DRO 使用的是更软的 advantage 估计形式,相关 advantage 需要在客户端侧构造后传入。
其实现方式为:
# Compute quadratic penalty term
quadratic_term = (target_logprobs - sampling_logprobs) ** 2
# Compute DRO objective
dro_objective = target_logprobs * advantages - 0.5 * beta * quadratic_term
# DRO loss is negative of objective
loss = -dro_objective.sum()dro 损失需要 Datum 的 loss_fn_inputs 中传入以下字段:
target_tokens: array[(N,), int]:target token IDs(来自采样器 )logprobs: array[(N,), float]:token 的sampling_logprobsadvantages: array[(N,), float]:RL 的优势值
输出:
logprobs: array[(N,), float]:token 的target_logprobs
指标:
loss_sum:SDK 返回的聚合损失,是一个标量
可以通过 loss_fn_config 自定义 :
fwd_bwd_future = await training_client.forward_backward_async(
data=data,
loss_fn="dro",
loss_fn_config={"beta": 0.05}
)