指南

损失函数

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

在监督学习中,我们实现了标准的交叉熵损失(即负对数似然),该损失优化策略 pθp_\theta 以最大化 token xx 的对数概率:

L(θ)=Ex[logpθ(x)]L(\theta) = -\mathbb{E}_x[\log p_\theta(x)]

其中 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()  # scalar

cross_entropy损失需要Datumloss_fn_inputs中传入target_tokensweights两个字段:

  • target_tokens: array[(N,), int] | array[(N, K), int]:target token IDs
  • weights: array[(N,), float] | array[(N, K), float]:token级的损失权重(通常来自渲染器)

输出:

  • logprobs: array[(N,), float] | array[(N, K), float]:请求的target token的对数概率

指标:

  • loss_sum:SDK 返回的聚合损失,是一个标量

importance_sampling

对于强化学习,我们实现了策略梯度目标的一个常见变体,适用于学习策略 pp 与采样策略 qq 存在差异的实际场景(例如由于非确定性导致的 off-policy 情况)。

问题在于,若两者存在差异,则目标:

L(θ)=Expθ[A(x)]L(\theta) = \mathbb{E}_{x \sim p_\theta}[A(x)]

由于 xqx \sim q(采样器)并不严格等同于期望的 xpθx \sim p_\theta(学习器),会导致估计有偏。为修正此偏差,我们采用改进的"重要性采样"目标:

LIS(θ)=Exq[pθ(x)q(x)A(x)]L_{\text{IS}}(\theta) = \mathbb{E}_{x \sim q}\left[\frac{p_\theta(x)}{q(x)} A(x)\right]

该目标可得到正确的期望奖励。公式中:

  • logpθ(x)\log p_\theta(x)target_logprobs)来自学习器,在 forward_backward 的前向阶段计算。
  • logq(x)\log q(x)sampling_logprobs)来自采样器,在采样时记录,用作修正项。

其实现方式为:

# Compute probability ratio
prob_ratio = torch.exp(target_logprobs - sampling_logprobs)
# Compute importance-weighted loss
loss = -(prob_ratio * advantages).sum()

importance_sampling 损失需要 Datumloss_fn_inputs 中传入以下字段:

  • target_tokens: array[(N,), int]:target token IDs(来自采样器 qq
  • logprobs: array[(N,), float]:token 的 sampling_logprobs
  • advantages: array[(N,), float]:RL 的优势值(正值表示强化,负值表示抑制)

输出:

  • logprobs: array[(N,), float]:token 的 target_logprobs

指标:

  • loss_sum:SDK 返回的聚合损失,是一个标量

ppo

PPO(Schulman et al., 2017)通过引入裁剪目标函数来解决标准策略梯度方法的问题,将策略更新限制在采样分布的邻域内,从而防止在同一 rollout 分布上进行多步梯度更新时出现过大的策略偏移。

该目标函数通过裁剪重要性比值 pθ(x)q(x)\frac{p_\theta(x)}{q(x)} 来防止策略更新幅度过大,其中 pθp_\theta 为学习器策略,qq 为采样策略。注意,PPO 的裁剪与损失计算均以 token 为单位独立进行。

PPO 裁剪目标为:

LCLIP(θ)=Exq[clip ⁣(pθ(x)q(x),1ϵlow,1+ϵhigh)A(x)]L_{\text{CLIP}}(\theta) = -\mathbb{E}_{x \sim q}\left[\text{clip}\!\left(\frac{p_\theta(x)}{q(x)},\, 1 - \epsilon_{\text{low}},\, 1 + \epsilon_{\text{high}}\right) \cdot A(x)\right]

最终 PPO 损失结合了裁剪与未裁剪两个目标:

LPPO(θ)=Exq[min ⁣(pθ(x)q(x)A(x),  clip ⁣(pθ(x)q(x),1ϵlow,1+ϵhigh)A(x))]L_{\text{PPO}}(\theta) = -\mathbb{E}_{x \sim q}\left[\min\!\left(\frac{p_\theta(x)}{q(x)} \cdot A(x),\; \text{clip}\!\left(\frac{p_\theta(x)}{q(x)},\, 1 - \epsilon_{\text{low}},\, 1 + \epsilon_{\text{high}}\right) \cdot A(x)\right)\right]

其中 ϵlow\epsilon_{\text{low}}ϵhigh\epsilon_{\text{high}} 为超参数(当前在 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 损失需要 Datumloss_fn_inputs 中传入以下字段:

  • target_tokens: array[(N,), int]:target token IDs(来自采样器 qq
  • logprobs: array[(N,), float]:token 的 sampling_logprobs
  • advantages: 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_future

cispo

CISPO(Clipped Importance Sampling Policy Optimization)是一种策略梯度方法。它与 PPO 都会使用重要性比值 pθ(x)q(x)\frac{p_\theta(x)}{q(x)},但差异在于:PPO 裁剪的是目标函数本身,而 CISPO 会先裁剪重要性比值,并将裁剪后的比值作为 target_logprobs 的梯度系数。

CISPO 目标为:

LCISPO(θ)=Exq[sg(clip(pθ(x)q(x),1ϵlow,1+ϵhigh))logpθ(x)A(x)]L_{\text{CISPO}}(\theta) = \mathbb{E}_{x \sim q}\left[\operatorname{sg}\left(\text{clip}\left(\frac{p_\theta(x)}{q(x)}, 1-\epsilon_{\text{low}}, 1+\epsilon_{\text{high}}\right)\right) \cdot \log p_\theta(x) \cdot A(x)\right]

其中 sg\operatorname{sg} 表示 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 损失需要 Datumloss_fn_inputs 中传入以下字段:

  • target_tokens: array[(N,), int]:target token IDs(来自采样器 qq
  • logprobs: array[(N,), float]:token 的 sampling_logprobs
  • advantages: array[(N,), float]:RL 的优势值

输出:

  • logprobs: array[(N,), float]:token 的 target_logprobs

指标:

  • loss_sum:SDK 返回的聚合损失,是一个标量

CISPO 默认使用单侧裁剪:不启用下界,只限制上界,即 clip_low_threshold=0.0clip_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 训练中,采样策略 qq 可能滞后于当前学习策略 pθp_\theta。如果设置较高的下界,会让已经被当前策略降低概率的旧 token 仍然获得较大的梯度系数,削弱重要性采样对 stale token 的衰减效果。因此,默认关闭下界通常更稳健。

dro

DRO(Direct Reward Optimization)是一种通用的 off-policy 甚至 offline 强化学习方法。它在奖励加权的 logprob 项之外加入二次惩罚,约束当前学习策略 pθp_\theta 不要相对采样策略 qq 偏移过大。

DRO 目标为:

LDRO(θ)=Exq[logpθ(x)A(x)12β(logpθ(x)q(x))2]L_{\text{DRO}}(\theta) = \mathbb{E}_{x \sim q}\left[\log p_\theta(x) \cdot A(x) - \frac{1}{2}\beta\left(\log \frac{p_\theta(x)}{q(x)}\right)^2\right]

其中 β\beta 控制二次惩罚强度。注意,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 损失需要 Datumloss_fn_inputs 中传入以下字段:

  • target_tokens: array[(N,), int]:target token IDs(来自采样器 qq
  • logprobs: array[(N,), float]:token 的 sampling_logprobs
  • advantages: array[(N,), float]:RL 的优势值

输出:

  • logprobs: array[(N,), float]:token 的 target_logprobs

指标:

  • loss_sum:SDK 返回的聚合损失,是一个标量

可以通过 loss_fn_config 自定义 β\beta

fwd_bwd_future = await training_client.forward_backward_async(
    data=data,
    loss_fn="dro",
    loss_fn_config={"beta": 0.05}
)
这篇文档对你有帮助吗?

本页目录