新闻详情

新闻详情

首页 / 资讯中心 / 详情

1 个脚本、13B 模型、19 倍提速:DeepSpeed-Chat 端到端 RLHF 训练完整指南

发布时间:2026/9/6 23:20:16来源:尧图网络
1 个脚本、13B 模型、19 倍提速:DeepSpeed-Chat 端到端 RLHF 训练完整指南
1 个脚本、13B 模型、19 倍提速DeepSpeed-Chat 端到端 RLHF 训练完整指南【免费下载链接】DeepSpeedDeepSpeed is a deep learning optimization library that makes distributed training and inference easy, efficient, and effective.项目地址: https://gitcode.com/GitHub_Trending/de/DeepSpeed普通微调框架为什么跑不动 RLHF一句话概括RLHF 不是换个损失函数而是一种推理与训练必须交替执行的负载——每一步迭代里模型先要逐 token 地生成回答内存带宽受限的解码再拿这些回答去算 PPO 梯度计算受限的训练。常规 DDP/微调框架没有为生成这个环节做任何加速结果就是整台机器长期趴在峰值性能的 5% 以下。DeepSpeed-Chat 给出的答案是 Hybrid EngineDeepSpeed-HE一个把训练引擎和推理引擎缝进同一个模型对象的统一系统配合对齐 InstructGPT 的三阶段流水线和单脚本入口让 OPT-13B 级别的 RLHF 训练在单节点上半天内完成、在 Azure 上花费不到 300 美元数据出处见 blogs/deepspeed-chat/README.md。先看核心机制Hybrid Engine 凭什么同时快又省这一节回答怎么做到的。结论先行同一个 PyTorch 模型对象通过eval()/train()两个标准接口在两套完全不同的前向路径之间切换——推理时走 DeepSpeed 推理内核 张量并行 KV-Cache 工作区训练时走 ZeRO 分片 LoRA 的常规路径。整个过程对用户透明你甚至不用换模型、不用搬权重。直观地说这就像同一间工厂白天跑重卡物流、晚上切换成精密装配线厂房模型权重不动换的是产线工艺和仓储方式。双引擎如何共用一个对象实现入口是 deepspeed/runtime/hybrid_engine.py 中的DeepSpeedHybridEngine它直接继承标准DeepSpeedEngine。构造时做三件事遍历模型按inference_policies把匹配的层Transformer、nn.Linear、nn.Embedding、nn.LayerNorm、OPT 位置编码逐一替换成推理容器_inference_containers同时把原始模块和原始forward存进_orig_modules/_orig_fwds备查若配置里inference_tp_size 1提前按 TP 规模切好通信组mp_group和ReplaceWithTensorSlicingout_dim0, in_dim1为推理侧的张量并行铺路若开启enable_cuda_graph用validate_cuda_graph_support校验 ZeRO 阶段是否兼容兼容则构建DecodeGraphCache。切换动作就藏在两个最不起眼的 API 里eval()推理/生成模式把每层的forward指到推理容器的 kernel 上并调用transform_for_inference()train()训练模式调用transform_for_training()后把forward指回原始实现恢复训练语义。step()里还有一步reset_params()训练步之后把推理容器的权重与最新分片同步保证下一轮生成用的是最新参数。换句话说无缝切换不是黑魔法而是一套精心设计的 forward 指针换装。生成侧KV-Cache、高性能内核与张量并行生成入口是generate()。它的工作流程可以拆成四步记录全局 batch 规模_total_batch_size并开始计时内部区分_gather_latency参数汇聚与_generate_latency真正解码两个阶段若启用了release_inference_cache先通过retake_inference_cache()重新向WorkspaceOp申请推理工作区——失败会gc.collect()empty_cache()后重试仍失败则抛RuntimeError执行_generate生成 token结束时若开启该开关调用workspace.release_workspace()把显存整块归还给训练侧。这套用完即还的显存策略正是博客所说在每种模式下重配置显存系统以最大化可用显存的具体落地生成阶段独享全部剩余显存做 KV-Cache训练阶段独享全部剩余显存放分片和梯度。两个值得注意的边界设计无匹配策略时的兜底populate_all_inference_policies()发现模型类型没有任何匹配的推理策略时会打印警告并把inference_policies置空generate()回退到模型原生generate()路径——不加速但保证功能可用CUDA Graph 的约束见 deepspeed/runtime/hybrid_engine_graph.pyZeRO-3、release_inference_cache、inference_tp_size 1三种情况都会禁用图缓存并警告。原因是 ZeRO-3 下推理容器不持有常驻权重图 replay 会固化过期的参数缓冲。另外DecodeGraphCache按解码位置逐个建图同一条序列内图回放与急切实行严禁混用避免序列长度计数器失步导致 KV-Cache 被悄悄写坏。训练侧ZeRO 与 LoRA 的可组合叠加训练阶段没有神秘的新机制关键在兼容性设计ZeRO 分片和 LoRA 适配器被实现为互相兼容、可任意组合的优化统一引擎下直接叠加。源码里generate()对 LoRA 的处理最能说明问题场景行为非 ZeRO-3生成前fuse_lora_weight()把 LoRA 权重融进推理容器生成后unfuse_lora_weight()还原ZeRO-3 pin_parametersGatheredParameters一次性 gather 全部非驻留参数若同时inference_tp_size 1则按tp_gather_partition_size默认 8 层一组分组 gather每组内apply_tensor_parallelism完成 TP 推理ZeRO-3 未 pin 参数生成中逐层_zero3_forward包裹 gather结束后用unfuse_lora_weight_non_pinned()在非驻留参数上完成 un-fuse也就是说训练时权重按 ZeRO 分片摊在各卡上生成时按 TP 方式重新组织切分结束后再切回——同一份权重两种切法双向无损往返。配置速查hybrid_engine 配置块deepspeed/runtime/config.py 中的HybridEngineConfig定义了写入 DeepSpeed JSON 的hybrid_engine块的全部字段与默认值字段类型 / 默认值作用enabledbool /False总开关启用 Hybrid Enginemax_out_tokensint /512生成最大长度同时作为推理容器的max/min_out_tokens传入也限制 CUDA Graph 捕获的解码位置数inference_tp_sizeint /1推理张量并行规模 1时启用mp_group与分区 gather 逻辑release_inference_cachebool /False生成后释放推理工作区、训练前重新申请压低显存峰值pin_parametersbool /TrueZeRO-3 下映射为gather_all_layers生成前把全部非 TP 层参数 gather 进显存常驻tp_gather_partition_sizeint /8ZeRO-3 TP 推理时分层 gather 的每组层数enable_cuda_graphbool /False启用 decode 阶段 CUDA Graph 缓存与 ZeRO-3 / TP / 工作区释放互斥一个能直接跑的最小样例在 tests/hybrid_engine/hybrid_engine_config.jsontrain_batch_size: 32、train_micro_batch_size_per_gpu: 2、zero_optimization.stage: 0且offload_param.device: cpu、stage3_param_persistence_threshold: 0、fp16.enabled: true、loss_scale_window: 100、gradient_clipping: 1.0。配套的 tests/hybrid_engine/hybrid_engine_test.py 用enable_hybrid_engineTrue初始化 OPT-350M依次执行m.eval()推理前向与m.train()训练前向正是上面双模式切换的最小验证。从 0 到跑通一条动线走完 RLHF 全流程机制讲完来看实际怎么用。下面的叙事顺序是装环境 → 一键三阶段 → 对话验证 → 换规模换集群 → 用 API 定制自己的策略。装环境与一键三阶段DeepSpeed-Chat 的训练示例在 DeepSpeedExamples 仓库的applications/DeepSpeed-Chat应用目录里。以 OPT-13B 当 actor、OPT-350M 当奖励模型为例整个安装到启动流程如下# 安装 DeepSpeedRLHF 流程要求 0.9.0 及以上版本 pip install deepspeed0.9.0 # 获取训练示例仓库DeepSpeedExamples 应用目录 git clone https://gitcode.com/GitHub_Trending/de/DeepSpeed cd DeepSpeedExamples/applications/DeepSpeed-Chat/ pip install -r requirements.txt # 单脚本跑完 SFT 奖励模型 RLHF 三个阶段约半天出 checkpoint python train.py --actor-model facebook/opt-13b --reward-model facebook/opt-350m --deployment-type single_node三阶段本身严格对齐 InstructGPTStep 1 用人工精选的 query-answer 对做监督微调SFTStep 2 训练一个通常更小的奖励模型 RW数据集是同一 query 多个答案的人工排序Step 3 用 PPO 以 RW 反馈继续微调 SFT 模型。流水线还有两个 InstructGPT 的可选项被内置为默认能力EMA 指数移动平均 checkpoint评测时可选用以及 Mixture Training把下一词预测目标混入 PPO 目标防止 SQuAD2.0 等基准能力退化。此外抽象数据集层统一不同数据格式切分/混合能力把多个数据集合理拌开后再分配到三个阶段——想喂多数据源的话这两项就是入口。三阶段耗时分解8× A100-40G 单 DGX 节点模型组合Step 1 (SFT)Step 2 (奖励模型)Step 3 (RLHF)总计Actor: OPT-13B, Reward: OPT-350M2.5h0.25h10.8h13.6h解读RLHF 阶段的耗时是前两步之和的十倍量级这也正是为什么整篇文章的机制部分都在围着生成加速转——省下的每一秒都花在 Step 3。用推理 API 验收模型训练完成后不需要自己搭推理栈DeepSpeed-Chat 自带对话式推理 API可以直接玩多轮交互官方示例里模型能这样回答Human: Do you know Microsoft? Assistant: Yes, Im a big fan of Microsoft. What about them? Human: Can you explain it to a 6-year old child? I wonder how I should describe it Assistant: Sure, I can try. Microsoft is a company that makes computers, and they make a program called Windows ...换规模、换集群只改两个参数同一条命令改--actor-model与--deployment-type就能覆盖从消费级单卡到 64 卡集群的所有形态# 66B 大模型64 卡 / 8 个 DGX 节点每节点 8× A100-80G约 9 小时跑完 python train.py --actor-model facebook/opt-66b --reward-model facebook/opt-350m --deployment-type multi_node # 1.3B 试跑单张消费级 GPU午餐时间约 2.2 小时拿到可玩的 checkpoint python train.py --actor-model facebook/opt-1.3b --reward-model facebook/opt-350m --deployment-type single_gpu模型组合Step 1Step 2Step 3总计硬件OPT-66B OPT-350M82 min5 min7.5h~9h64× A100-80G8 节点OPT-1.3B OPT-350M2900s670s1.2h~2.2h单张 A6000 48G解读从 1.3B 到 66B 只需要换模型名和部署类型脚本内部按deployment-type自动选择对应的 ZeRO/并行配置——这是单脚本体验的核心卖点。用 RLHF API 定制自己的策略不想被固定流程绑住的话API 层把一次 RLHF 迭代拆成两个显式动作generate_experience走 Hybrid Engine 的推理加速路径产出经验和train_rlhf用 PPO 目标更新 actor 与 critic# 构建统一的 RLHF 引擎actor critic 两个模型tokenizer 与总迭代数一并传入 engine DeepSpeedRLHFEngine( actor_model_name_or_pathargs.actor_model_name_or_path, critic_model_name_or_pathargs.critic_model_name_or_path, tokenizertokenizer, num_total_itersnum_total_iters, argsargs) # PPO 训练器承接引擎推理/训练两阶段在 trainer 上显式分离 trainer DeepSpeedPPOTrainer(engineengine, argsargs) for prompt_batch in prompt_train_dataloader: # 阶段一基于提示词生成经验Hybrid Engine 推理加速路径 out trainer.generate_experience(prompt_batch) # 阶段二以 PPO 目标同时更新 actor 与 critic奖励模型 actor_loss, critic_loss trainer.train_rlhf(out)这个先经验、后更新的两段式抽象就是前面讲的推理/训练双阶段在 API 层的投影——换算法时你只需要重排这两个调用的方式底层的显存管理与切换逻辑不用碰。数据说话横向对比与扩展拐点横向对比线与 Colossal-AI、原生 PyTorch 驱动的 HuggingFace 方案相比DeepSpeed-RLHF 的单卡 Step 3 吞吐提升超过 10 倍上图为单张 A100-40G 对比缺图标的条目代表对方 OOM多卡端8× A100-40G 单节点、不同模型规模的端到端吞吐相对 Colossal-AI 加速 6–19 倍相对 HuggingFace DDP 加速 1.4–10.5 倍。上图为 1.3B 模型单次 RLHF 迭代的时间分解大部分时间落在生成阶段。DeepSpeed-HE 靠高性能推理内核把这一阶段相对 HuggingFace 提速最高 9 倍、相对 Colossal-AI 提速 15 倍——端到端效率的差距基本全部来自这里。模型规模上限同样拉开明显差距同样硬件下 Colossal-AI 单卡最大 1.3B、单节点 A100-40G 最大 6.7B而 DeepSpeed-HE 分别可达 6.5B 与 50B最大可训规模扩大至 7.5 倍。扩展拐点线先说清结构。本流水线中生成阶段约占 20% 的计算量、训练阶段占 80%但生成阶段要对 256 token 提示逐 token 解码 256 token是内存带宽受限的环节实际可能吃掉端到端时间的大头训练阶段则是对 512 token 完整样本做少量前向/反向的计算受限环节吞吐表现稳定。为此 DeepSpeed-HE 两手抓两阶段都用尽量大的 batch生成阶段模型装得下单卡就用高性能内核榨满带宽装不下时改用张量并行而非 ZeRO 扩卡——TP 的 GPU 间通信更少带宽利用率不掉。上图显示6.7B–66B 区间效率最高上探到 175B 后因显存限制撑不住大 batch吞吐回落但仍比 1.3B 模型高效 1.2 倍且加卡加大 batch 后还有提升空间。其有效性能比对比系统高 19 倍——反过来看那些系统实际只跑在峰值 5% 以下。上图是 13B/66B actor350M reward随 DGX 节点数增加的扩展曲线整体在 64 GPU 内扩展良好但形状分两段——小规模超线性大规模近线性或次线性。原因在于 ZeRO 分片与全局 batch 上限的相互作用卡越多模型状态摊得越薄单卡显存越富余能塞更大的单卡 batch于是超线性但扩到一定规模后最大全局 batch本场景为 1024 组 query-answer 对、序列长 512给单卡 batch 封顶只能近线性甚至次线性。推论很实用给定全局 batch 上限最优吞吐与成本效率恰好落在超线性与次线性的交界处该拐点主要由单卡可跑的最大 batch可用显存 × 全局 batch 的函数决定——选卡数时照着这个拐点选即可。选型与边界决策参考 把分散在各处的数字汇总成决策表。单卡能训多大Hybrid Engine 显存管理下单卡规格V100 32GA6000 48GA100 40GA100 80G最大可训模型OPT-2.7BOPT-6.7BOPT-6.7BOPT-13B解读单卡即可训练 13B 模型没有多卡资源的团队也能产出可用模型这是该方案区别于必须集群的常规 RLHF 工具的关键卖点。集群端成本参考均为 Step 3 时长规模单节点 8× A100-80G多节点 64× A100-80GOPT-13B9h约 $2901.25h约 $320OPT-30B18h约 $5804h约 $1024OPT-66B2.1 天约 $16207.5h约 $1920OPT-175B—20h约 $5120解读13B/30B 在单节点内最划算66B 以上明显是多节点集群的主场175B 一天内可收工。对比数据的适用前提引用前必读以上所有数字针对Step 3基于 DeepSpeed-RLHF 精选数据集与固定配方——135M tokens 训一个 epoch67.5M query tokens131.9k 条 query、长度 25667.5M 生成 tokens131.9k 条回答、长度 256每步最大全局 batch 为 0.5M tokens1024 组 query-answer 对。换数据集或 batch 规格后耗时与成本都需重新评估。什么场景该选它目标是端到端对齐 InstructGPT 的 RLHF 全流程、希望单脚本交付、且模型规模从单卡 13B 到集群 175B 之间任取时Hybrid Engine 是当前把生成加速 ZeRO/LoRA 训练捏成一个系统的完整方案。反过来若你的负载是纯推理服务或纯 SFT 微调直接用对应的 DeepSpeed 推理/训练引擎即可不必引入 RLHF 全流程。小结与关键入口清单 DeepSpeed-Chat 用三件事解决了 RLHF难上手、太贵、扩不动单脚本三阶段训练加推理 API 的低门槛体验内置 EMA 与混合训练、支持多数据源混合的 InstructGPT 完整流水线以及把训练与推理统一进 Hybrid Engine 的系统设计——同一个模型对象在 ZeRO 分片训练与 TP 推理内核之间随train()/eval()自由换装。引用该项目建议使用 arXiv:2308.01320BibTeX 见 blogs/deepspeed-chat/README.md。关键源码/配置/测试入口混合引擎核心generate()、LoRA 融合/还原、ZeRO-3 分区 gather、workspace 回收deepspeed/runtime/hybrid_engine.pyCUDA Graph 解码缓存按位置建图、序列级启用决策、ZeRO-3/TP 互斥校验deepspeed/runtime/hybrid_engine_graph.pyhybrid_engine配置块定义7 个字段与默认值deepspeed/runtime/config.py端到端测试与最小配置样例tests/hybrid_engine/性能数据与流水线说明原始出处blogs/deepspeed-chat/README.md需要说明成本与耗时数据基于发布时的训练配方135M tokens、单 epoch、每步 0.5M tokens 全局 batch对比时请以其基准规格为准仓库当前实现已在此基础上演进CUDA Graph、共享 prefill 工作区等具体行为以源码与测试为准。【免费下载链接】DeepSpeedDeepSpeed is a deep learning optimization library that makes distributed training and inference easy, efficient, and effective.项目地址: https://gitcode.com/GitHub_Trending/de/DeepSpeed创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

闭环温度控制系统设计全流程:从传感器选型到PID调参实战 2026/9/6 23:56:21

闭环温度控制系统设计全流程:从传感器选型到PID调参实战

简介:一份完整的闭环温度控制系统电子工程设计报告,面向自动化与电子信息类学生及单片机开发者,适用于课程设计、毕业设计或项目预研中的温度控制方案构思。报告以8051单片机为核心,系统阐述了从功能指标制定、方案选型到硬件电路…

阅读更多 →
基于深度学习YOLOv8+PyQt5的车牌检测识别系统实战解析 2026/9/6 23:56:21

基于深度学习YOLOv8+PyQt5的车牌检测识别系统实战解析

直接开门见山这次要拆解的项目是“基于深度学习 YOLOv8 PyQt5 的车牌检测识别系统”。这是一个非常典型的计算机视觉桌面应用:前端用 PyQt5 做界面,后端用 YOLOv8 做目标检测,再配合 OCR 技术完成车牌字符识别。整个项目覆盖了目标检测、字符…

阅读更多 →
7.20 OLED的OLED_Refresh函数栈溢出踩内存 2026/9/6 23:56:21

7.20 OLED的OLED_Refresh函数栈溢出踩内存

void OLEDS_Refresh(void) {u8 i, n;//这里可能会导致栈溢出 换成静态变量staticstatic uint8_t buf[129]; /* buf[0]控制码 0x40, buf[1..128]像素数据 */for (i 0; i < 8; i){/* 设置页地址和列地址 */OLEDS_WR_Byte(0xB0 i, OLEDS_CMD); /* 页地址 0~7 */OLEDS_WR_B…

阅读更多 →
LangChain核心实战:从大模型调用到Agent开发与LangGraph选型 2026/9/6 23:56:21

LangChain核心实战:从大模型调用到Agent开发与LangGraph选型

刚开始接触 LangChain 的时候&#xff0c;很多同学都有这样的困惑&#xff1a;今天看到一个教程在讲 Model&#xff0c;明天又看到一个教程在讲 Agent&#xff0c;后天又冒出 LangGraph、RAG、CrewAI。各种概念堆在一起&#xff0c;感觉每个单词都认识&#xff0c;但就是串不起…

阅读更多 →
LangChain与Agent工程实践:从模型调用到稳定可控的大模型应用流程编排 2026/9/6 23:56:21

LangChain与Agent工程实践:从模型调用到稳定可控的大模型应用流程编排

我最早有“LangChain 到底有没有用”这个概念&#xff0c;是在一个做知识库问答的周会上。当时团队里的模型接口已经跑通&#xff0c;单条提问和回答都很正常&#xff0c;可一旦放进一个带历史记录、带文档检索、带多轮追问的真实产品里&#xff0c;整个流程就开始失控&#xf…

阅读更多 →
基于Spring Boot的自助健身房管理系统的设计与实现 2026/9/6 23:53:21

基于Spring Boot的自助健身房管理系统的设计与实现

摘 要 本毕业设计旨在开发一套基于Spring Boot框架的自助健身房管理系统&#xff0c;解决传统健身房人工依赖重、效率低等痛点。系统采用B/S架构&#xff0c;后端利用Java语言与MySQL数据库实现业务逻辑持久化&#xff0c;前端通过Vue框架与微信小程序提供双向交互&#xff0…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞