当前位置: 首页 > news >正文

如何训练奖励模型:train-llm-from-scratch的Bradley-Terry损失详解

如何训练奖励模型:train-llm-from-scratch的Bradley-Terry损失详解

【免费下载链接】train-llm-from-scratchA straightforward method for training your LLM, from downloading data to generating text.项目地址: https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch

train-llm-from-scratch是一个从零训练大语言模型(LLM)的开源项目,覆盖数据下载、预训练、SFT 到 RLHF 全流程。其中奖励模型(Reward Model)是 RLHF 的关键一环:它给模型回答打一个分数,分数越高代表越受人类偏好。本文将带你用Bradley-Terry 损失亲手训练一个奖励模型,并讲清每一步背后的原理。

奖励模型在整体流程中的位置

奖励模型不是孤立存在的,它处于整个后训练流水线的中心位置——上游接收 SFT 模型,下游为 PPO 提供奖励信号:

从图中可以看到:预训练基座 → SFT 指令微调之后,分三条支路:

  • Reward Model(Bradley-Terry):训练出奖励模型,为 PPO 提供reward signal
  • DPO / ORPO / KTO:不依赖奖励模型的偏好对齐路线;
  • GRPO / RLVR:使用可验证奖励的强化学习路线。

三条路最终都会汇入评估与聊天环节。

奖励模型长什么样:给 Transformer 加一个"打分头"

奖励模型的核心思路很简单(沿用 InstructGPT 的做法):

  1. 取 SFT 训练好的 Transformer 主干,扔掉语言模型头(lm_head)
  2. 在最上面接一个极小的标量奖励头Linear(n_embed → 1),把隐藏状态压成一个数;
  3. 奖励值从最后一个真实 token的隐藏状态读取。

为什么是最后一个 token?因为注意力是因果的(causal),序列末尾的 token "看过"了整段内容,却不会注意到它右侧的填充(padding)——所以连注意力掩码都省了。

模型定义见 RewardModel:

源码中的注释写得很直白:a scalar reward head on top of a Transformer backbone(在 Transformer 主干上叠加标量奖励头)。

Bradley-Terry 损失详解:整个训练信号就一行代码

给定同一个 prompt 下的两个回答:chosen(偏好)rejected(拒绝),奖励模型要学的事情只有一件——让 chosen 的分数高于 rejected。Bradley-Terry 模型把这变成如下损失:

L = -log σ(r_chosen - r_rejected)

其中 σ 是 sigmoid 函数。直观理解:

情况r_chosen − r_rejected损失 L含义
分数差距很大(判对)≫ 0≈ 0几乎没有梯度,模型"满意"
两个分数相同(瞎猜)00.693随机水平,训练起点
分数判反了≪ 0很大强梯度,逼模型纠正

整个损失函数只有三行(bradley_terry_loss):

def bradley_terry_loss(chosen_rewards, rejected_rewards): return -F.logsigmoid(chosen_rewards - rejected_rewards).mean()

除了损失,还有两个配套指标(reward_train.py):

  • preference_accuracy(偏好准确率)r_chosen > r_rejected的样本占比,这是最值得盯的核心指标
  • reward_margin(奖励间隔):平均r_chosen − r_rejected,衡量模型是否还在有效区分好坏回答。

偏好数据从哪来:一条完整的数据管道

Bradley-Terry 训练需要{"prompt", "chosen", "rejected"}三元组。项目的数据管道把公共数据集(HH-RLHFUltraFeedback)转换成标准 JSONL 格式:

对应脚本是 prepare_preference_data.py,它会:

  • 分别从 HH-RLHF(人类标注的 helpful/harmless 偏好)和 UltraFeedback(LLM 评审的偏好对)拉取数据;
  • 过滤空 prompt、chosen == rejected 的噪声样本;
  • 输出训练集preferences.jsonl留出测试集preferences_test.jsonl(偏好准确率就在它上面测)。

批次构造逻辑在 preference_dataset.py:chosen 和 rejected 两侧各自经 chat template 渲染成 token,右侧填充(right-padding)到同长——这在因果注意力下是安全的。

实战:一条命令训练你的奖励模型

先确认已完成前序阶段(预训练 → SFT,产出sft.pt)和偏好数据,然后运行:

# 单卡 PYTHONPATH=. python scripts/train_reward.py # 双卡并行 PYTHONPATH=. torchrun --standalone --nproc_per_node=2 scripts/train_reward.py

训练入口 train_reward.py 做了这些事:

  1. 从 SFT 检查点加载主干sft_ckpt),加上奖励头;
  2. 每个 batch 把 chosen 与 rejected拼成 2B 条序列,一次前向拿到所有奖励值,再对半拆分;
  3. 计算 Bradley-Terry 损失 → 反向传播 → AdamW + 余弦学习率更新参数;
  4. 定期在测试集上评估偏好准确率并保存检查点。

关键超参在 configs/reward.json 中,默认配置一目了然:

参数默认值说明
batch_size8每步 8 对偏好(实际过 16 条序列)
lr1e-5精细微调级别的学习率
max_len768单条最大 token 长度
grad_clip1.0梯度裁剪保稳定

想调参可以直接传命令行参数,例如--lr 1e-5 --max_len 768

训练结果怎么看:三个数字背后的含义

训练日志会打印loss / train_acc / test_acc / margin,官方文档 docs/04_reward_model.md 给出了清晰的解读:

  • loss:从 0.693(纯随机)起步,随分数差距拉大而下降;
  • train_acc / test_acc:在干净的小数据上能到 1.0;在真实的、带噪声的HH-RLHF / UltraFeedback 上,0.65–0.75 才是正常水平——人类偏好本身就很嘈杂,别为 0.7 焦虑;
  • margin:还在持续增长说明模型仍在有效区分好坏回答。

训练完成后,奖励模型保存为reward.pt,供下游 PPO 直接加载。

下一步:奖励模型如何被 PPO 消费

PPO 用奖励模型给每次采样(rollout)的回答打分,作为强化学习的目标信号:

在 train_ppo.py 中通过--reward_source rm即可切换为使用训练好的奖励模型(代码内部调用 load_reward_model 加载)。

📌小贴士:如果你的任务不需要通用偏好打分(比如数学题有标准答案),也可以跳过奖励模型,直接走 DPO 这条"免奖励模型"路线:

相关实现见 dpo.py 与 train_dpo.py。

总结:5 个文件看懂奖励模型训练

文件职责
src/post_training/reward_model.py奖励模型:SFT 主干 + 标量奖励头
src/post_training/reward_train.pyBradley-Terry 损失与评估指标
data_loader/preference_dataset.py偏好数据批次构造(右侧填充)
scripts/train_reward.py训练入口:单卡 / 双卡 DDP
configs/reward.json默认超参配置

一句话回顾:给 SFT 模型加一个打分头 → 用 Bradley-Terry 损失在偏好对上训练 → 偏好准确率 0.65–0.75 即为合格 → 交给 PPO 当"老师"。跟着 docs/04_reward_model.md 与 POST_TRAINING.md,你就能完整复现这一经典 RLHF 环节。

【免费下载链接】train-llm-from-scratchA straightforward method for training your LLM, from downloading data to generating text.项目地址: https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

http://www.cnnetsun.cn/news/4350519.html

相关文章:

  • 我的世界Overlay测试指南:拼好种与终末之诗通关验证
  • 暴跌行情下,用市场温度判断短线与长线交易逻辑
  • JobOps签证赞助商查询教程:一站式验证UK签证担保公司资质
  • 上运动神经元 vs 下运动神经元:从解剖到瘫痪定位诊断一次讲清
  • 从采集设置到可视化流程:用清源AI搭建游戏调度决策助手
  • 小智音箱蓝牙通信实战:ESP32+SPP透传与调试全攻略
  • 基于Django的智能图书管理系统:数据驱动、分析与推荐一体化实践
  • Loop Engineering深度解析:闭环原理、核心要素与工程落地
  • 交银金科后端岗笔试复盘:考点、编程题与避坑指南
  • PageIndex 自托管部署:三步在本地搭好无向量文档索引
  • SpringBoot+Vue人事管理系统:从源码拆解到实战部署
  • Next AI Draw.io 部署指南:10 分钟跑通 AI 画图
  • Next AI Draw.io:一句话画出 draw.io 图表,从 Docker 部署到模型选型的完整上手指南
  • 用RTX 4090打造AI魔镜:本地大模型与多模态视觉实战
  • Nessus安装与使用教程
  • OpenVoice语音克隆实操指南:3分钟搭好环境,5秒语音样本完成克隆
  • 写论文英文AI率太高,怎么降低?先校对时态,再重组固定句式。
  • 如何实现天猫多店防关联管理自动化?无人值守订单处理,日发5000单零差错
  • 麻将实战总打错?从牌效率到防守,拆解“一看就会,一打就费”的真相
  • AI歌声生成全流程:从本地部署到未修音干声处理
  • Halo邮箱验证:注册即发验证码,把假邮箱挡在门外
  • IP地址与二进制转换全解析:从手算方法到Python实现
  • 为什么DNSHE免费DNS解析这么快?Anycast DNS技术原理深度剖析
  • 从零搭建RAG知识库问答系统:原理、代码与工程落地
  • Frigate 完整安装教程:30 分钟部署本地监控 AI NVR 与实时对象检测
  • 用Jetpack Compose从零实现安卓计时器:状态驱动UI与协程实战
  • AI搜索,你的企业还在“隐形”?北京GEO优化公司推荐,这三家助你提升可见度
  • CCR 配置备份与恢复的完整实战笔记
  • 如何为pm-skills贡献一个新技能?完整开发流程与验证脚本使用指南
  • 手术场景视觉-轨迹联合预测模型:从原理到工程部署