Skip to content

损失函数

Weaver 的训练 API 通过 loss_fn 字符串选择服务端 loss,并通过 Datum.loss_fn_inputs 传入与 token 对齐的张量。不同 loss 的输入约定不同,但核心原则一致:model_input 提供上下文,target_tokens 提供 next-token 目标,其他张量控制每个 token 的训练权重或策略梯度。

常用 Loss

loss 名称典型用途主要输入
cross_entropySFT、蒸馏、带 mask 的 next-token 训练target_tokensweights
forward_logprob只取 logprobs,不反传target_tokens,可选 sampling_mask
importance_samplingRL policy gradient、离线/在线采样修正target_tokenslogprobsadvantagesloss_mask
grpoGRPO 风格训练importance_sampling 类似,由 trainer/recipe 组织 group advantage
ppo_clipPPO clipped surrogatetarget_tokenslogprobsadvantagesloss_mask
truncated_importance_sampling带截断的 IS 训练importance_sampling 类似
opd_kl_importance_sampling带 teacher KL 路径的 OPD/IS 训练old_logprobsteacher_logprobs 等额外输入
surrogateSDK 自定义 loss 的内部 backward 通道surrogate_weightsloss_mask

TIP

大多数用户从 cross_entropyimportance_sampling 开始即可。更高级的 PPO/GRPO/OPD loss 通常由 NexRL recipe 或自定义 trainer 生成对应字段。

Cross Entropy

cross_entropy 是标准监督微调 loss,适合 SFT、格式学习、任务蒸馏和小规模数据验证。

输入字段

每条 Datum 需要:

  • target_tokensTensor[int64],与 model_input 做 next-token 对齐。
  • weightsTensor[float32],每个 target token 的 loss 权重。
python
datum = types.Datum(
    model_input=types.ModelInput.from_ints(tokens[:-1]),
    loss_fn_inputs={
        "target_tokens": torch.tensor(tokens[1:], dtype=torch.int64),
        "weights": torch.tensor(weights[1:], dtype=torch.float32),
    },
)

计算方式

其中:

  • 是第 个目标 token。
  • 是该 token 的权重。
  • 是模型在前文条件下预测该 token 的概率。

Mask 策略

最常见的 SFT mask 是忽略 prompt,只训练 completion:

python
weights = [0.0] * len(prompt_tokens) + [1.0] * len(completion_tokens)

如果某些答案片段更重要,也可以使用大于 1.0 的权重。确保 weights[1:]target_tokens 等长。

forward_logprob

forward_logprob 用于只计算 token logprobs,不做反向传播。它常用于:

  • 评估模型对答案的 likelihood。
  • 计算 perplexity。
  • 为自定义 loss 准备 logprob 张量。
  • 在 RL 中保存采样策略的 old logprobs。
python
result = training_client.forward(
    datums,
    "forward_logprob",
    wait=True,
)

SamplingClient.compute_logprobs() 不同,trainer 侧 forward_logprob 对齐的是显式 target_tokens,不会在返回列表前添加首 token 的 None 占位。

Importance Sampling

importance_sampling 面向策略优化。它使用 rollout 时的 logprobs 和 advantage,训练当前策略提高高 advantage token 的概率,并降低低 advantage token 的概率。

输入字段

每条 Datum 通常需要:

  • target_tokens:目标 token。
  • logprobs:rollout 或旧策略下的 token logprobs。
  • advantages:与 target_tokens 对齐的 advantage/reward signal。
  • loss_mask:0/1 mask,控制哪些 token 参与 RL loss。

可选字段:

  • ref_logprobs:参考策略 logprobs,用于 KL 正则。
  • sampling_mask:采样 mask,常用于结构化采样或 router replay 场景。
python
datum = types.Datum(
    model_input=types.ModelInput.from_ints(input_tokens),
    loss_fn_inputs={
        "target_tokens": torch.tensor(target_tokens, dtype=torch.int64),
        "logprobs": torch.tensor(old_logprobs, dtype=torch.float32),
        "advantages": torch.tensor(advantages, dtype=torch.float32),
        "loss_mask": torch.tensor(loss_mask, dtype=torch.int64),
    },
)

直觉公式

简化写法如下:

其中:

  • loss_mask
  • 是 advantage。
  • 是当前策略与采样策略的概率比修正。

实际实现还会根据 loss_fn_config 处理 ratio 变换、截断、KL 正则和聚合方式。

PPO / GRPO 相关 Loss

Weaver trainer 中提供了 grpoppo_cliptruncated_importance_samplingopd_kl_importance_sampling 等 loss 入口。它们主要给 NexRL 或自定义 trainer 使用,因为这些算法通常需要 group-level reward、old policy logprobs、reference policy logprobs、teacher logprobs 或额外的 KL 配置。

建议做法:

  • 如果你使用 NexRL,优先让 NexRL recipe/trainer 生成对应 Datum 字段。
  • 如果你自己写 trainer,确保所有 per-token 张量都与 target_tokens 等长。
  • 先用 wait=True 跑小 batch 验证 loss 输出和 metrics,再切换异步训练。

自定义 Loss

如果内置 loss 不满足需求,可以在本地定义 PyTorch loss,并使用 forward_backward_custom()

python
def ranking_loss(data, logprob_tensors):
    chosen, rejected = logprob_tensors
    loss = -(chosen.sum() - rejected.sum()).sigmoid().log()
    metrics = {"ranking_loss": float(loss.detach())}
    return loss, metrics


result = training_client.forward_backward_custom(datums, ranking_loss)
training_client.optim_step(types.AdamParams(), wait=True)

这个 API 的内部流程是:

  1. Weaver 执行 forward(..., "forward_logprob")
  2. SDK 把返回的 logprobs 转成带梯度的 PyTorch 张量。
  3. 你的函数返回标量 loss 和 metrics。
  4. SDK 对该 loss 反传,并把 logprob 梯度作为 surrogate loss 送回 Weaver。

调试建议

  • 先确认 target_tokensweightsadvantagesloss_mask 的长度一致。
  • 对 prompt token 使用 0 mask 或 0.0 weight,避免训练模型复读输入。
  • RL loss 中的 logprobs 应来自采样/旧策略,而不是当前训练 step 之后的模型。
  • 对新 loss 先用很小 batch 跑 forward(),检查返回 logprobs 和 metrics。

下一步

Weaver API 中文文档