损失函数
Weaver 的训练 API 通过 loss_fn 字符串选择服务端 loss,并通过 Datum.loss_fn_inputs 传入与 token 对齐的张量。不同 loss 的输入约定不同,但核心原则一致:model_input 提供上下文,target_tokens 提供 next-token 目标,其他张量控制每个 token 的训练权重或策略梯度。
常用 Loss
| loss 名称 | 典型用途 | 主要输入 |
|---|---|---|
cross_entropy | SFT、蒸馏、带 mask 的 next-token 训练 | target_tokens、weights |
forward_logprob | 只取 logprobs,不反传 | target_tokens,可选 sampling_mask |
importance_sampling | RL policy gradient、离线/在线采样修正 | target_tokens、logprobs、advantages、loss_mask |
grpo | GRPO 风格训练 | 与 importance_sampling 类似,由 trainer/recipe 组织 group advantage |
ppo_clip | PPO clipped surrogate | target_tokens、logprobs、advantages、loss_mask |
truncated_importance_sampling | 带截断的 IS 训练 | 与 importance_sampling 类似 |
opd_kl_importance_sampling | 带 teacher KL 路径的 OPD/IS 训练 | old_logprobs、teacher_logprobs 等额外输入 |
surrogate | SDK 自定义 loss 的内部 backward 通道 | surrogate_weights、loss_mask |
TIP
大多数用户从 cross_entropy 和 importance_sampling 开始即可。更高级的 PPO/GRPO/OPD loss 通常由 NexRL recipe 或自定义 trainer 生成对应字段。
Cross Entropy
cross_entropy 是标准监督微调 loss,适合 SFT、格式学习、任务蒸馏和小规模数据验证。
输入字段
每条 Datum 需要:
target_tokens:Tensor[int64],与model_input做 next-token 对齐。weights:Tensor[float32],每个 target token 的 loss 权重。
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:
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。
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 场景。
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 中提供了 grpo、ppo_clip、truncated_importance_sampling 和 opd_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():
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 的内部流程是:
- Weaver 执行
forward(..., "forward_logprob")。 - SDK 把返回的 logprobs 转成带梯度的 PyTorch 张量。
- 你的函数返回标量 loss 和 metrics。
- SDK 对该 loss 反传,并把 logprob 梯度作为
surrogateloss 送回 Weaver。
调试建议
- 先确认
target_tokens、weights、advantages、loss_mask的长度一致。 - 对 prompt token 使用
0mask 或0.0weight,避免训练模型复读输入。 - RL loss 中的
logprobs应来自采样/旧策略,而不是当前训练 step 之后的模型。 - 对新 loss 先用很小 batch 跑
forward(),检查返回 logprobs 和 metrics。