ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

veRL SPIN 在线自我博弈训练实战:环境搭建、Online DPO 核心算法与踩坑记录

veRL SPIN 在线自我博弈训练实战:环境搭建、Online DPO 核心算法与踩坑记录 文档教程人工智能大模型RLHF【免费下载链接】Awesome-ML-SYS-TutorialMy learning notes for ML SYS.项目地址https://gitcode.com/gh_mirrors/aw/Awesome-ML-SYS-Tutorial点击查看免费下载本文以仓库 rlhf/verl/spin/test-log.md 为主线结合同目录下的 SPIN-dev.md 与 dev-log.md 展开。文章完整还原了在 veRL 中运行 SPINSelf-Play Fine-tuNing自我博弈微调其工程实现等价于 Online DPO的端到端流程——从 SGLang 容器构建、分支化 verl 安装、数据与模型准备到recipe/spin/test.sh的实际启动同时深入 Online DPO 的损失函数、训练五阶段与 veRL 源码级改造点main_dpo、RayDPOTrainer、core_algos损失函数、update_policy_dpo等并记录ref_update_freq导致 reward hacking 的关键踩坑。读完本文读者可以完整复现一次基于 SGLang veRL 的 SPIN 训练并理解其背后的算法与工程实现。一、SPIN 是什么从自我博弈到 Online DPOSPINSelf-Play Fine-tuNing的核心思想是让策略模型与自身的副本进行对弈上一轮迭代的策略π_t充当对手当前模型π_θ需要学会区分自己生成的回复与人类偏好数据中的回复从而在迭代中持续逼近目标分布。在实际工程实现中SPIN 的更新目标与 Online DPO 高度一致——不再依赖离线静态偏好对而是在线采样策略模型的输出并用奖励模型Reward Model或判断器Judge即时判定被选中/被拒绝回复然后套用 DPO 风格的对比损失进行更新。1.1 Online DPO 的核心组件根据 SPIN-dev.md 的梳理Online DPO 需要四个核心组件Policy Model策略模型被训练的模型即当前迭代的π_θReference Model参考模型固定的基准模型通常是 Policy Model 的冻结副本用于提供对数概率的参照系评估组件二选一Reward Model奖励模型评分模型为每个生成结果打分Judge判断器比较器比较两个生成结果并选出更好的一个。1.2 核心损失函数Online DPO 的核心是最大化被选中回复相对被拒绝回复的概率比。在 veRL 的 SPIN 实现中dev-log.md 明确指出在core_algos模块中实现了两种损失变体Sigmoid 损失经典 DPO$$\mathcal{L}{\text{DPO}}(\theta) -\mathbb{E}{(x, y_w, y_l) \sim \mathcal{D}}\left[\log \sigma\left(\beta \cdot \left(\log \frac{p_\theta(y_w|x)}{p_{\text{ref}}(y_w|x)} - \log \frac{p_\theta(y_l|x)}{p_{\text{ref}}(y_l|x)}\right)\right)\right]$$IPO 损失基于平方差$$\mathcal{L}{\text{IPO}}(\theta) \mathbb{E}{(x, y_w, y_l) \sim \mathcal{D}}\left[\left(\log \frac{p_\theta(y_w|x)/p_{\text{ref}}(y_w|x)}{p_\theta(y_l|x)/p_{\text{ref}}(y_l|x)} - \frac{1}{2\beta}\right)^2\right]$$其中各符号含义如下符号含义$p_\theta(yx)$策略模型对给定提示 $x$ 生成回复 $y$ 的概率$p_{\text{ref}}(yx)$参考模型冻结的策略模型副本的相应概率$\beta$控制 KL 约束强度的超参数调控模型更新的幅度$y_w$ / $y_l$被选中的回复 / 被拒绝的回复两个损失函数都通过超参数beta调控更新幅度并在返回时输出损失均值。若读者希望进一步对照 SPOSelf-Play Preference Optimization原始论文的损失形式及其在实际训练中采用的近似公式可参考同目录系列文档 rlhf/verl/sppo/paper.md。二、环境搭建构建 SGLang 容器SPIN 训练依赖 SGLang 作为 actor 的 rollout 引擎因此第一步是进入lmsysorg/sglang:latest容器。test-log 中给出了完整的容器启动命令docker run -it --name xxx --gpus all \ --shm-size32g \ --ipchost \ -v /root/.cache:/root/.cache \ -e HF_TOKENxx \ lmsysorg/sglang:latest \ /bin/bash各参数的作用如下--gpus all暴露全部 GPU供多卡 rollout 与训练使用--shm-size32g扩大共享内存。多进程/多节点通信如 NCCL、DataLoader 的 shared memory对/dev/shm容量敏感过小的 shm 会导致通信初始化失败--ipchost与宿主机共享 IPC 命名空间避免容器内进程间通信的隔离限制-v /root/.cache:/root/.cache挂载 Hugging Face 缓存目录复用已下载的模型权重与数据集缓存避免重复下载-e HF_TOKENxx注入 Hugging Face 访问令牌用于拉取受控/私有模型实际使用时替换为真实 token 值lmsysorg/sglang:latest官方 SGLang 镜像内置 CUDA 环境与推理引擎依赖。说明SPIN 训练场景下actor 的 rollout 由 SGLang 承担而训练阶段由 veRL 的训练引擎完成这正是 veRLhybrid engine设计将 actor 的 rollout engine 与 training engine 放在同一资源组串行执行的典型应用背景可参见 rlhf/verl/readme.md。三、安装 verlspin 分支进入容器后按以下步骤安装带 SPIN 改造的 verl。注意 test-log 中安装的是作者分支cedricbeta/verl的spin分支这是该测试日志对应的实际代码来源mkdir -p /tmp chmod 1777 /tmp apt update apt install -y python3.10 python3.10-venv python3 -m ensurepip --upgrade sudo apt install tmux python3 -m venv ~/.python/sglang source ~/.python/sglang/bin/activate python3 -m pip install uv python3 -m uv pip install wheel python3 -m uv pip install packaging python3 -m uv pip install flash-attn --no-build-isolation --no-deps cd ~ git clone https://github.com/cedricbeta/verl.git cd verl git checkout spin python3 -m uv pip install -e .[sglang]分步要点chmod 1777 /tmp修正容器内/tmp的写权限sticky bit避免某些构建/缓存流程因权限问题失败安装python3.10及 venv 模块并用ensurepip补齐 piptmux用于在容器内维持长时间训练会话防止 SSH 断开导致训练中断在 SGLang 镜像的既有 Python 环境中新建独立 venv~/.python/sglang隔离依赖使用uvpython3 -m pip install uvpython3 -m uv pip install ...加速包安装flash-attn采用--no-build-isolation --no-deps安装直接复用容器预编译环境避免重编译 FlashAttention 的漫长等待最后git checkout spin切换到 SPIN 实现分支并以可编辑模式-e .[sglang]安装 verl 及其 SGLang 推理依赖便于后续调试源码。四、wandb 登录与数据、模型准备训练过程中需要日志与可视化先登录 wandbwandb login随后准备数据集与基础模型。test-log 选用GSM8K 数学推理数据集与Qwen2.5-3B-Instruct 基座模型注意SPIN 训练是在该 Instruct 模型之上继续做偏好对齐这与 SPPO 系列文档中直接以 Qwen2.5-7B-Instruct 起步的做法一致可见 rlhf/verl/sppo/compare_with_ppo_grpo.mdpython3 examples/data_preprocess/gsm8k.py --local_dir ~/data/math huggingface-cli download Qwen/Qwen2.5-3B-Instruct --local-dir $HOME/models/Qwen2.5-3B-Instructexamples/data_preprocess/gsm8k.py是 verl 仓库内置的数据预处理脚本将 GSM8K 转为 RLHF 训练所需的 parquet 格式并输出到~/data/math训练集train.parquet与测试集test.parquethuggingface-cli download ... --local-dir将模型权重下载到本地目录训练脚本中通过actor_rollout_ref.model.path直接引用该路径。五、启动训练recipe/spin 下的 test.sh数据与模型就绪后执行训练test-log 注明已在H20 x44 卡 H20上实测通过export CUDA_VISIBLE_DEVICES0,1,2,3 cd recipe/spin bash test.shCUDA_VISIBLE_DEVICES0,1,2,3将 4 张卡暴露给训练进程与trainer.n_gpus_per_node4的资源设定对应recipe/spin/test.sh位于 verl 仓库内是 SPIN 分支随附的训练脚本其内部通常以python3 -m verl.trainer.main_dpo启动在线 DPO 训练并会通过trainer.ref_update_freq控制参考模型/对手策略的刷新频率详见下文踩坑记录。可对照 PPO-RM 与 GRPO 的同类 recipe 脚本rlhf/verl/sppo/compare_with_ppo_grpo.md理解其配置结构。六、从测试日志到源码实现veRL 的 SPIN 改造点test-log 仅是运行记录真正支撑这条命令链的是 verl spin 分支中的三处核心改造。根据 dev-log.md 的实现日志可以还原如下6.1 新增 main_dpo 入口与 RayDPOTrainermain_dpo作为在线 DPO 的入口脚本负责加载配置、初始化 Ray 集群、构造 DPO 训练器并启动训练流程。它基于原 PPO 入口main_ppo改造将更新阶段切换到 DPO 流程RayDPOTrainer重用 PPO 的资源池管理、worker 分组与数据加载逻辑仅在训练更新阶段调用新的 DPO 更新接口利用core_algos中实现的对比损失sigmoid 或 IPO直接更新策略模型。6.2 core_algos 中的损失计算在核心算法模块中新增两个损失函数compute_online_dpo_losssigmoid 版本利用策略模型与参考模型对数概率比的差值计算基于 sigmoid 的 DPO 损失IPO 版本基于平方差形式计算损失。两者均以beta为超参数调控更新幅度返回损失均值——这与第一节给出的两个损失公式一一对应。6.3 PPO Worker 的 DPO 补丁升级在DataParallelPPOActor中新增update_policy_dpo方法与传统 PPO 更新步骤类似但它接收通过 union 合并的 chosen / rejected 回复从meta_info中提取chosen_mask随后调用核心算法模块中的 DPO 损失函数计算损失并执行反向传播与梯度更新在ActorRolloutRefWorker中新增update_actor_dpo方法为update_policy_dpo提供上层接口。从源码结构看SPIN 分支的工程策略是尽量复用 PPO 的分布式训练骨架Ray 资源池 DataParallelWorker只替换损失函数与更新入口这与RayDPOTrainer重用 PPO 资源池管理、只换更新逻辑的描述一致。七、Online DPO 训练五阶段详解SPIN-dev.md 以 TRL 的 OnlineDPO 为参照拆解了完整的训练流程正好与上述 veRL 改造点互相印证。五个阶段如下阶段一生成rollout为每个提示生成两个不同的回复采样两次# 为每个提示生成两个不同的回复 prompts inputs[prompt] # 形状: [batch_size] batch_size len(prompts) # 使用vLLM或标准生成 if use_vllm: prompt_ids, prompt_mask, completion_ids, completion_mask _generate_vllm(model, prompts) else: prompt_ids, prompt_mask, completion_ids, completion_mask _generate(model, prompts)这一阶段从输入批次提取提示、为每个提示采样两个不同回复并检查哪些回复包含 EOS 结束标记供后续未收敛回复处理使用。在 veRL SGLang 的组合中这里的生成由 SGLang rollout engine 完成。阶段二计算模型概率# 计算策略模型的对数概率 logprobs _forward(model, prompt_ids, prompt_mask, completion_ids, completion_mask) # 计算参考模型的对数概率(无梯度) with torch.no_grad(): if ref_model is not None: ref_logprobs _forward(ref_model, prompt_ids, prompt_mask, completion_ids, completion_mask) else: # PEFT情况只需禁用adapter with model.disable_adapter(): ref_logprobs _forward(model, prompt_ids, prompt_mask, completion_ids, completion_mask)要点策略模型带梯度计算对数概率参考模型在torch.no_grad()下冻结计算若使用 PEFTLoRA 等可通过禁用 adapter 的方式复用同一份权重得到参考概率——这与 veRL 中update_policy_dpo接收 union 合并的 chosen/rejected 回复后一次前向得到两侧 logprob 的设计思路一致。阶段三评估生成结果# 解码生成的回复 completions processing_class.batch_decode(completion_ids, skip_special_tokensTrue) if judge is not None: # 使用判断器进行对比评估 ranks judge.judge(prompts, list(zip(completions[:batch_size], completions[batch_size:]))) mask torch.tensor([rank 0 for rank in ranks], devicedevice) else: # 使用奖励模型进行评分 scores reward_model(prompt_completion_ids).scores # 处理未包含EOS的回复可选降低它们的分数 if missing_eos_penalty is not None: scores[~contain_eos_token] - missing_eos_penalty # 分割分数并比较 first_half, second_half scores.split(batch_size) mask first_half second_half这一阶段将 token ID 解码回文本用 Judge 或 Reward Model 对每对回复判定优劣确定被选中与被拒绝对未包含 EOS 的回复可施加missing_eos_penalty惩罚。阶段四组织数据并计算损失# 获取被选中和被拒绝回复的索引 batch_range torch.arange(batch_size, devicedevice) chosen_indices batch_range (~mask * batch_size) rejected_indices batch_range (mask * batch_size) # 获取被选中和被拒绝回复的对数概率 chosen_logprobs_sum, rejected_logprobs_sum torch.split(cr_logprobs_sum, batch_size) chosen_ref_logprobs_sum, rejected_ref_logprobs_sum torch.split(cr_ref_logprobs_sum, batch_size) # 计算对数概率比值 pi_logratios chosen_logprobs_sum - rejected_logprobs_sum ref_logratios chosen_ref_logprobs_sum - rejected_ref_logprobs_sum # 计算DPO损失所需的logits logits pi_logratios - ref_logratios # 根据指定的损失类型计算损失 if loss_type sigmoid: losses -F.logsigmoid(beta * logits) elif loss_type ipo: losses (logits - 1 / (2 * beta)) ** 2 loss losses.mean()这一阶段通过 mask 区分 chosen/rejected计算策略与参考模型之间的对数概率比差值logits再按loss_type选择 sigmoid 或 IPO 损失——这正是core_algos中compute_online_dpo_loss的对应逻辑也是第一节两个公式的代码形态。阶段五更新模型# 执行反向传播 if n_gpu 1: loss loss.mean() # 多GPU上平均损失 accelerator.backward(loss, **kwargs) # 返回损失 return loss.detach() / args.gradient_accumulation_steps执行反向传播并交给优化器更新参数多 GPU 下先对 loss 取均值返回时除以梯度累积步数以对齐累积语义。八、踩坑记录ref_update_freq与 Reward Hackingdev-log.md 记录了一个关键教训trainer.ref_update_freq设置过大。在 SPIN / Online DPO 中参考模型对手策略π_t需要定期从当前策略刷新以维持自我博弈的有效性。若ref_update_freq过大参考模型长期停留在旧的策略状态与当前策略的差距被拉大模型很容易被 reward hacking钻奖励模型的空子陷入局部最优local max最终导致训练崩盘。作者在重新调整该参数后训练能够稳定涨点并收敛。这也是所有在线/自博弈类算法SPIN、SPO、SPPO共同需要注意的节奏问题——参考策略的刷新频率直接决定了博弈对手的强度与训练稳定性。相关迭代式损失的设计可进一步对照 rlhf/verl/sppo/paper.md 中 SPO 的更新公式。九、小结与参考文档本文以 rlhf/verl/spin/test-log.md 为骨架完整覆盖了从容器构建、分支安装、数据准备到 H20x4 实测运行的 SPIN 训练链路并借助 SPIN-dev.md 与 dev-log.md 展开算法与工程细节Online DPO 的 sigmoid / IPO 损失、五阶段训练流程、veRL 的main_dpo/RayDPOTrainer/core_algos/update_policy_dpo改造以及ref_update_freq的调参陷阱。进一步阅读SPIN 训练算法详解Online DPO 组件、损失公式与五阶段伪代码SPIN 实现日志main_dpo、RayDPOTrainer、core_algos 与 worker 改造记录SPO 论文损失解析SPIN 系列原始目标与实用近似PPO-RM 与 GRPO recipe 对照同一训练框架下的其他算法配置veRL 框架解析hybrid engine、single/multi controller 等背景概念。赞分享文档教程人工智能大模型RLHF【免费下载链接】Awesome-ML-SYS-TutorialMy learning notes for ML SYS.项目地址https://gitcode.com/gh_mirrors/aw/Awesome-ML-SYS-Tutorial点击查看免费下载相关推荐TinyZero 算法扩展实战基于 verl single_controller 三步实现 Online DPO 等新 RL 算法TinyZero 算法扩展实战基于 verl single_controller 三步实现 Online DPO 等新 RL 算法 veRL本仓库 Tiny人工智能大模型强化学习推理模型pyrefly-numpy-stubs 实战指南为 Pyrefly 构建带数组形状信息的 NumPy 类型桩pyrefly numpy stubs 实战指南为 Pyrefly 构建带数组形状信息的 NumPy 类型桩 本篇技术指南围绕 Pyrefly 开源仓库中的文档教程人工智能大模型RLHFDouZero代码剖析深入理解模型训练与自我对弈的底层实现DouZero是一个基于深度强化学习技术的斗地主AI系统采用自我对弈机制从零开始训练。这个开源项目在ICML 2021上发表通过创新的Deep Monte人工智能强化学习深度学习AI 应用游戏开发上一篇OpenShift Origin 的 Kubernetes Rebase 全流程指南基于 openshift/kubernetes fork 的版本升级实践下一篇Dart Style常见问题解答解决格式化错误、样式冲突的实用方案创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表