ARTICLE DETAIL

资讯详情

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

Open-Assistant 模型复现全指南:基于 OASST 数据从 SFT、RM 到 RLHF 的完整训练流水线

Open-Assistant 模型复现全指南:基于 OASST 数据从 SFT、RM 到 RLHF 的完整训练流水线 Open-Assistant 模型复现全指南基于 OASST 数据从 SFT、RM 到 RLHF 的完整训练流水线【免费下载链接】Open-AssistantOpenAssistant is a chat-based assistant that understands tasks, can interact with third-party systems, and retrieve information dynamically to do so.项目地址: https://gitcode.com/gh_mirrors/op/Open-Assistant本文以 model/README.md 为骨架完整梳理 Open-Assistant 在model/目录下提供的最小化复现命令从环境初始化、OASST 数据准备到监督微调SFT、奖励模型RM与强化学习RL/PPO三阶段训练并深入仓库源码讲解配置解析、数据过滤与消息格式的底层实现。读完本文你将能依据仓库内真实的 YAML 配置与 Python 源码从零跑通一条可复用的对话模型训练流水线并理解每一步背后的数据流与代码依据。一、整体流水线三阶段复现路线Open-Assistant 的模型训练遵循经典的 RLHF人类反馈强化学习三阶段范式对应 model/model_training 目录下的三个独立训练入口阶段训练入口配置文件目标产物1. SFT 监督微调trainer_sft.pymodel/model_training/configs/config.yaml具备对话能力的基线策略模型2. RM 奖励模型trainer_rm.pymodel/model_training/configs/config_rm.yaml对回复质量打分排序的奖励模型3. RL 强化学习trainer_rl.pymodel/model_training/configs/config_rl.yaml经 PPO 优化后的最终策略模型model/README.md 给出的是最少命令跑通全流程的最小化复现指引。正式训练前请确认Python 版本不低于 3.10——原文档特别提示更低版本会触发 typer 工具的已知兼容性问题issue 编号 371 所描述的现象。二、环境准备数据目录与共享模块2.1 初始化数据与模型目录所有训练产物都有固定的落盘位置。原文档的第一步是创建目录并导出环境变量mkdir -p .cache mkdir -p .saved_models export DATA_PATH$PWD/.cache export MODEL_PATH$PWD/.saved_modelsDATA_PATH即配置中的cache_dir用于存放 OASST 消息树 JSONL 文件与 HuggingFace 数据集缓存MODEL_PATH用于存放 SFT、RM、RL 三阶段产出的模型目录。2.2 安装训练依赖与 OASST 数据模块除根目录的 Python 工程外训练代码还依赖oasst_data数据读写模块。根据 model/model_training/README.md安装方式为# 安装训练工程pyproject.toml 位于 model/ 目录 pip install -e .. # 安装 oasst 数据模块 python -m pip install ../../oasst-data/oasst_data的源码位于 oasst-data/oasst_data训练端正是通过它读取导出的消息树在 oasst_dataset.py 中read_message_trees(input_file_path)本地 JSONL与read_dataset_message_trees(hf_dataset_name, splittrainvalidation)HuggingFace 数据集是两条等价的数据入口。安装后可运行测试验证环境pytest .。若tests/test_patched_gpt_neox.py::test_flash_attention_patch抛出SystemExit按警告安装flash_attn即可python -m pip install flash_attn三、数据准备配置 OASST 训练数据源3.1 两种数据源配置方式在 model/model_training/configs/config.yamlSFT、config_rm.yamlRM或 config_rl.yamlRL中新建或修改一个配置节通过datasets下的oasst_export项指定数据方式 A本地 OASST JSONL 文件支持.jsonl与.jsonl.gz用input_file_path指定文件可置于cache_dir即DATA_PATH下也可使用绝对路径cp /path/to/oasst.trees.jsonl $DATA_PATHmy_data_config: datasets: - oasst_export: input_file_path: oasst_export.trees.jsonl.gz方式 BHuggingFace 数据集用hf_dataset_name指定数据集名my_data_config: datasets: - oasst_export: hf_dataset_name: OpenAssistant/oasst1优先级规则原文档明确若同时指定hf_dataset_name与input_file_pathinput_file_path优先生效。这一逻辑在 oasst_dataset.py 的load_oasst_export中有直接体现先判断input_file_path再回退到hf_dataset_name二者都未指定则抛出RuntimeError。3.2 oasst_export 的常用可选参数源码层面load_oasst_export还支持以下常用参数均可在oasst_export:下配置参数默认值作用langen逗号分隔的 ISO-639-1 语言码白名单只有tree.prompt.lang命中才保留仓库实际配置常用bg,ca,cs,da,de,en,es,fr,hr,hu,it,nl,pl,pt,ro,ru,sl,sr,sv,ukval_split0.2验证集划分比例仓库多数配置使用0.05top_kNone仅保留排名前 k 的 assistant 回复用于精选高质量对话input_max_length—限制输入序列长度以仓库内置的oasst_only配置节为例config.yamloasst_only: save_strategy: epoch datasets: - oasst_export: lang: bg,ca,cs,da,de,en,es,fr,hr,hu,it,nl,pl,pt,ro,ru,sl,sr,sv,uk hf_dataset_name: OpenAssistant/oasst1 #input_file_path: 2023-04-12_oasst_ready.trees.jsonl.gz #top_k: 1 val_split: 0.05 sort_by_length: false use_custom_sampler: false3.3 数据过滤逻辑树状态与角色约束oasst_dataset.py 展示了数据进入训练前的过滤管线这是理解同一份 OASST 数据如何服务三个不同阶段的关键树状态过滤SFT/RM 模式只接受tree_state ready_for_export的消息树RL 模式额外接受prompt_lottery_waiting状态等待提示词抽签的树脏数据过滤thread_filter剔除含已删除deleted或合成synthetic消息的线程top_k生效时按m.rank过滤超出阈值的回复阶段差异化切分SFT 模式只取以 assistant 回复结尾的完整线程若末尾是 prompter 叶子则pop()掉RM 模式以prompter 消息为前缀、多条带 rank 的回复为候选项构造对比样本RL 模式取所有以 prompter 结尾的前缀作为强化学习的状态输入。四、第一阶段SFT 监督微调4.1 启动命令cd model_training # 导出共享模块 export PYTHONPATH$PYTHONPATH:../../oasst-shared python trainer_sft.py --configs defaults oa_dataset_only pythia --cache_dir $DATA_PATH --output_dir $MODEL_PATH/sft_model # 如需使用 wandb 记录实验追加 --wandb_entity your_username/team_name两点提醒基于当前仓库源码--configs可以叠加多个配置节后出现的配置覆盖先出现的同名键最终合并结果再被命令行同名参数覆盖。该合并逻辑位于 trainer_sft.py 的argument_parsing先read_yamls(./configs)载入全部配置节再以defaults为底依次conf.update(...)示例中的oa_dataset_only/pythia是原文档写作时的占位配置名。当前仓库 config.yaml 中实际可用的等价配置包括oasst_only、oasst_export_eu仅 OA 数据以及pythia-70m-deduped、pythia-1B、pythia-6.9B、pythia-12B等模型配置节。若传入不存在的配置名程序会打印Could not find the config ...并退出。4.2 defaults 关键超参数解读defaults是 SFT 的基础配置config.yaml理解这些参数是调优的起点参数默认值说明learning_rate1e-5AdamW 学习率per_device_train_batch_size/per_device_eval_batch_size2/2单卡批大小受显存约束gradient_accumulation_steps32梯度累积步数等效全局批大小 单卡批 × 累积步 × GPU 数warmup_steps600学习率预热步数max_length/val_max_length512/ 空训练/验证最大序列长度空则沿用max_lengthnum_train_epochs3训练轮数eval_steps/save_steps200/1000评估与保存间隔save_strategy: stepssave_total_limit4最多保留的 checkpoint 数dtypefp16训练精度可选fp16/bf16/fp32等weight_decay0.00权重衰减max_grad_norm2.0梯度裁剪loss_fnCrossEntropyLoss损失函数poly_eps控制 poly loss 的平滑系数random_offset_probability0.8训练时随机截取消息片段的概率数据增强label_maskingtrue只对 assistant 回复计算损失的标签掩码use_system_prefix/use_system_tagfalse是否在开头注入系统前缀/系统标签peft_model/peft_typefalse/lora是否使用 LoRA 等参数高效微调这些超参在 trainer_sft.py 中被逐一映射到 HuggingFaceTrainingArgumentsdtype为fp16/float16时开启fp16为bf16/bfloat16时开启bf16。数据侧训练与评估分别构造DialogueDataCollator见 dialogue_collator.py其中random_offset_probability、label_masking、samples_mixing、系统前缀/标签等开关都会影响每条样本的最终拼装方式。4.3 更换底座模型要更换更大规模的 Pythia 模型可在config.yaml新建配置节或直接用--model_name覆盖为EleutherAI/pythia-{size}-dedupedmy-model: learning_rate: 8e-6 model_name: EleutherAI/pythia-6.9b-deduped weight_decay: 0.0 max_length: 2048 warmup_steps: 20 gradient_checkpointing: false gradient_accumulation_steps: 2 per_device_train_batch_size: 4 per_device_eval_batch_size: 4原文档特别强调模型越大通常越需要同步调低--learning_rate、调小--per_device_train_batch_size。此外若所选模型缺少pad_token、eos_token、sep_token需要修改 utils/utils.py 的get_tokenizer以注入正确的特殊 token仓库内置的tokenizer_sanity_check会在启动时打印特殊 token 及其 id便于核对。4.4 获取 SFT checkpoint 并测试# 选择指定 checkpoint export SFT_MODEL$MODEL_PATH/sft_model/checkpoint-X # 或自动取最新 checkpoint export SFT_MODEL$MODEL_PATH/sft_model/$(ls -t $MODEL_PATH/sft_model/ | head -n 1)训练完成后可用 tools/model_cli.py 交互式测试--8bit用于 8bit 量化模型python3 tools/model_cli.py --model_path saved_path/huggingface五、第二阶段奖励模型 RM 训练5.1 启动命令cd model_training python trainer_rm.py --configs defaults_rm oasst-rm-1-pythia-1b同样地oasst-rm-1-pythia-1b为原文档示例配置名当前仓库 config_rm.yaml 中实际提供的是oasst-rm-1-pythia-1.4b、oasst-rm-1-pythia-2.8b、oasst-rm-1-pythia-6.9b三个规模版本可按显存选择。由于模型配置节写得极简必须用defaults_rm补齐其余默认项原文档强调这一点。5.2 defaults_rm 关键参数参数默认值说明is_reward_modeltrue切换为奖励模型输出标量分数而非 token 概率poolinglast序列池化方式将最后一层 hidden state 聚合为奖励分loss_fnRMLoss成对排序损失score_l2_reg: 0.001给分数加 L2 正则metrics[accuracy, kendalltau]评估指标排序准确率与 Kendall tau 相关系数max_replies5每个提示最多使用的候选回复数warmup_steps/save_steps10/100RM 训练通常收敛更快间隔更短奖励模型的数据以对比对为主oasst-rm-1-pythia-2.8b的配置节config_rm.yaml展示了典型构成oasst_exportOA 排名数据augment_oasst增强 OA 数据anthropic_rlhf、shp、hellaswag、webgpt、hf_summary_pairs等外部对比数据集。源码层RM 数据经 ranking_collator.py 组织为前缀 多条回复的对比 batchRMTrainer.compute_losstrainer_rm.py直接对RMLoss求梯度。5.3 获取 RM checkpoint# 选择指定 checkpoint export REWARD_MODEL$MODEL_PATH/reward_model/checkpoint-X # 或自动取最新 checkpoint export REWARD_MODEL$MODEL_PATH/reward_model/$(ls -t $MODEL_PATH/reward_model/ | head -n 1)六、第三阶段RLHF / PPO 强化学习6.1 启动命令拿到 SFT 与 RM 模型后进入强化学习阶段cd model_training python trainer_rl.py --configs defaults_rlhf --cache_dir $DATA_PATH --rank_model $REWARD_MODEL --sft_model $SFT_MODEL --output_dir $MODEL_PATH/rl_modeldefaults_rlhf的关键参数config_rl.yaml参数默认值说明batch_size1PPO 每步采样的 batchchunk_size1经验 chunk 大小num_rollouts16每步 rollout 数量total_steps10000总训练步数eval_size/num_eval_prompts500/64评估集大小与评估提示数rank_config/sft_config空内嵌的 RM / SFT 子配置模型名、dtype、池化方式等以pythia_rlhf配置节config_rl.yaml为例它通过triton_host_rm/triton_host_sft指定 RM 与 SFT 模型的推理服务地址并在rank_config/sft_config中分别声明两套模型参数——RL 阶段将 SFT 作为待优化策略、RM 作为奖励来源通过 PPO 算法持续优化。6.2 基于 Triton 的多 GPU 推理部署进阶model/model_training/README.md 给出了在大规模训练时的完整部署路径适用于 8 卡 GPU 服务器# 1. 构建 Triton 容器需先安装 singularity singularity build --sandbox tritonserver-pyt.sif docker://nvcr.io/nvidia/tritonserver:22.08-pyt-python-py3 # 2. 将训练好的 RM / SFT 模型转换为 tritonserver 模型仓库 python to_triton.py --configs pythia_rlhf --triton_mode rm python to_triton.py --configs pythia_rlhf --triton_mode sft # 3. 分别在 GPU 7、GPU 6 上启动 RM 与 SFT 推理服务 SINGULARITYENV_CUDA_VISIBLE_DEVICES7 singularity run --nv --bind .triton_models/model_store_rm:/model_store tritonserver-pyt.sif tritonserver --model-repository/model_store --http-port 8001 --grpc-port 8002 --metrics-port 8003 SINGULARITYENV_CUDA_VISIBLE_DEVICES6 singularity run --nv --bind .triton_models/model_store_sft:/model_store tritonserver-pyt.sif tritonserver --model-repository/model_store --http-port 8004 --grpc-port 8005 --metrics-port 8006 # 4. 通过 accelerate 启动 PPO 训练 export TRITON_HOST_RMlocalhost:8002/RM_MODEL_NAME export TRITON_HOST_REFlocalhost:8005/REF_MODEL_NAME CUDA_VISIBLE_DEVICES0,1,2,3,4,5 OMP_NUM_THREADS1 accelerate launch --main_process_port 29501 --config_file configs/accelerate_config.yaml --num_processes 6 trainer_rl.py --configs defaults defaults_rlhf pythia_rlhf oasst_export_latin_cyrillic_rlhf注意--num_processes必须等于训练所用 GPU 数量上述命令为 6。多进程启动配置见 configs/accelerate_config.yamlPPO 相关参数见 configs/ppo_config.yaml。七、消息与 Token 格式模型理解对话的协议model/README.md 末尾指向 model/MESSAGE_AND_TOKEN_FORMAT.md这份配套文档是理解训练数据拼装规则的关键。7.1 Token 基础模型接收的输入并非字母文本而是被切分成 token 序列每个 token 有全局唯一的 id词表 vocab。查看模型附带 JSON 中的词表时会看到奇怪的Ġ字符这是bytes_to_unicode函数将 0–255 的码位中的控制字符与空白字符整体上移 2560x100以使其可打印所致——空格码位 320x20因此变成Ġ码位 2880x120。在 Open-Assistant 的对话格式中角色标签|prompter|、|assistant|、|system|也在词表中以特殊 token 形式存在定义见 formatting.py 的QA_SPECIAL_TOKENS。7.2 三阶段训练范式model/MESSAGE_AND_TOKEN_FORMAT.md 从训练原理角度重新定义了三个阶段先在大规模互联网文本上预训练得到基座 LLM如 GPT-3、Galactica 系列再进入SFT——用志愿者构建的 OASST 演示数据学习监督策略得到基线模型随后RM阶段由志愿者对 SFT 输出投票生成对比数据训练奖励模型最后PPO阶段用奖励模型进一步微调 SFT 模型产出最终的策略模型。其中 SFT 阶段可能混入其他数据集但每个数据集都必须被改造成仓库要求的统一对话格式。7.3 Message Format v2多数 Open-Assistant 模型使用|prompter|{prompt}|endoftext||assistant|格式没有独立的前缀标签前缀可放在第一条 prompter 消息之前或之内You are a large language model that wants to be helpful|prompter|Hello!|endoftext||assistant|模型回复时会在结尾补齐零个或多个|endoftext|这只是为了让 batch 内各样本等长。7.4 Message Format v2-new实验中的新格式新格式引入系统前缀与完整历史对话|system|{prefix}|endoftext|在prefix为空时整体省略以下换行与注释仅为可读性实际格式不含|system|{prefix}|endoftext| # 历史对话对逐条排列 |prompter|{history_prompt[i]}|endoftext||assistant|{history_reply[i]}|endoftext| # 当前待回复的新提示 |prompter|{prompt}|endoftext||assistant|示例换行仅用于展示|system|You are a large language model that wants to be helpful|system| |prompter|What is red and round?|endoftext||assistant|Hmm, a red balloon?|endoftext| |prompter|No, smaller|endoftext||assistant|7.5 源码中的格式实现v2-new 的角色标签QA_SPECIAL_TOKENS中|system|对应的format_system_prefix函数formatting.py负责拼装系统前缀系统标签属性注入Utterance.system_tagformatting.py会把lang、quality、humor、creativity、length、context等属性随机乱序写入系统标签其中property_dropout以一定概率丢弃属性、add_length控制是否附加估算的消息长度compute_length按单词数估算训练侧开关use_system_prefix/use_system_tag/system_property_dropout/system_add_length由 config.yaml 的defaults提供并被传入 trainer_sft.py 的DialogueDataCollator最终决定每条训练样本的拼装形态。仓库中的use_system_tag配置节展示了其开启方式system_add_length: True。八、多数据集混合与子采样除 OASST 数据外SFT 阶段还混合了大量外部指令数据默认datasets列表见 config.yaml包含webgpt、squad_v2、adversarial_qa、xsum、cnn_dailymail、soda、joke、gsm8k、多语种翻译对wmt2019_zh-en等。model/model_training/README.md 提供的数据集规模统计工具# 统计全部数据集规模注意会下载超过 100GB 的数据 python check_dataset_counts.py --datasets all --mode sft # 只统计指定数据集 python check_dataset_counts.py --datasets webgpt squad_v2 --mode sft子采样fraction 与 size通过fraction按比例随机抽取或size按条数随机抽取可对训练数据子采样注意数据集名后要加冒号datasets: - webgpt: fraction : 0.05 - prompt_dialogue: size : 500 - adversarial_qa - trivia_qa_nocontext上述配置下每个 epoch 会从webgpt随机抽取 5%、从prompt_dialogue随机抽取 500 条、未指定参数的adversarial_qa与trivia_qa_nocontext全量使用且每个 epoch 抽取的随机子集都不同。原文档确认该机制兼容torch.distributed分布式训练。多语种翻译数据集wmt2019_*、ted_trans_*通过名称后缀指定语言对例如wmt2019_zh-en、wmt2019_ru-en、ted_trans_nl-en当前仅支持ar,de,fr,en,it,nl,tr,ru,ms,ko,ja,zh这些语言的提示词翻译。九、训练优化与问题排查9.1 DeepSpeed 支持编辑 configs/zero_config.json当前默认 Zero 阶段 3可调整 ZeRO 策略仓库另提供zero3_config_sft.json、zero3_config_falcon.json、zero3_config_pretrain.json等针对不同模型与任务的配置。启用方式是在命令末尾追加--deepspeed一般推荐用 deepspeed 启动器deepspeed trainer_sft.py --configs defaults your-model-name --deepspeed9.2 断点续训在 trainer_sft.py 中通过--resume_from_checkpoint从最近保存的 checkpoint 恢复训练同时作用于 wandb 续写。为快速验证流程可使用仓库内置的debug配置节config.yaml选用pythia-70m-deduped小模型、fp32精度、关闭 wandb、每 20 步评估保存是理想的冒烟测试模板。9.3 环境问题多 GPU 训练若在 VM 上训练可能需要安装 OpenMPImpi4py安装需要python-devDebian/Ubuntu 下sudo apt install libpython3.10-dev版本号与本地 Python 对应flash_attnuse_flash_attention: true或相关测试需要安装flash_attnpip install flash_attn。十、关键文件速查用途路径复现总入口文档model/README.md消息与 Token 格式规范model/MESSAGE_AND_TOKEN_FORMAT.md训练工程详细说明model/model_training/README.mdSFT / RM / RL 配置config.yaml、config_rm.yaml、config_rl.yamlSFT / RM 训练入口trainer_sft.py、trainer_rm.pyOASST 数据加载与过滤custom_datasets/oasst_dataset.py对话格式拼装与特殊 tokencustom_datasets/formatting.py模型/分词器/数据集工厂utils/utils.pyTriton 模型转换to_triton.py交互式模型测试tools/model_cli.pyOASST 数据读写模块oasst-data/oasst_data综上所述Open-Assistant 在model/目录中提供了一条完全可复现的 RLHF 训练流水线以input_file_path/hf_dataset_name两种方式接入 OASST 数据经由统一的对话消息格式v2 / v2-new与DialogueDataCollator完成样本拼装依次通过trainer_sft.py、trainer_rm.py、trainer_rl.py三个入口产出 SFT、RM 与最终策略模型。读者可从oasst_onlypythia-70m-deduped的小规模组合起步逐步替换为更大的 Pythia、Llama 或 Falcon 底座并结合 DeepSpeed、LoRA 与 Triton 部署方案扩展到多卡环境。【免费下载链接】Open-AssistantOpenAssistant is a chat-based assistant that understands tasks, can interact with third-party systems, and retrieve information dynamically to do so.项目地址: https://gitcode.com/gh_mirrors/op/Open-Assistant创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表