ARTICLE DETAIL

资讯详情

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

TRL Callbacks 完全指南:Rich 进度条、Completions 日志、BEMA 权重平均与 Weave 追踪

TRL Callbacks 完全指南:Rich 进度条、Completions 日志、BEMA 权重平均与 Weave 追踪 TRL Callbacks 完全指南Rich 进度条、Completions 日志、BEMA 权重平均与 Weave 追踪【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trlTRLTransformer Reinforcement Learning基于 transformers 的TrainerCallback机制为 SFT、DPO、GRPO、KTO、RLOO 等强化学习/偏好对齐训练流程提供了五个开箱即用的回调RichProgressCallback终端可视化进度、LogCompletionsCallback将模型生成记录到 WB/Comet、BEMACallback偏差校正指数移动平均、WeaveCallbackWB Weave 追踪与评估以及SyncRefModelCallback训练中同步参考模型。本文以 docs/source/callbacks.md 为主线结合 trl/trainer/callbacks.py 的源码实现与 tests/test_callbacks.py、tests/test_rich_progress_callback.py 的测试用例逐一讲解每个回调的用途、参数语义、底层原理与接入方式帮助你为自己的训练实验挑选并组合合适的回调。回调机制速览TRL 如何复用 transformers 的 TrainerCallbackTRL 的全部回调都定义在 trl/trainer/callbacks.py 中并统一继承自transformers.TrainerCallback。这意味着它们拥有与 transformers 训练循环完全一致的生命周期钩子on_train_begin、on_step_end、on_log、on_evaluate、on_train_end等可以无缝挂载到 TRL 提供的任何 Trainer 上。所有回调都在 trl/trainer/init.py 与 trl/init.py 中被导出因此可以直接从顶层导入from trl import BEMACallback, LogCompletionsCallback, RichProgressCallback, SyncRefModelCallback, WeaveCallback挂载方式有两种等价于 transformers 的标准用法# 方式一构造 Trainer 时通过 callbacks 参数传入 trainer DPOTrainer(..., callbacks[LogCompletionsCallback(trainertrainer)]) # 方式二训练前动态添加 trainer DPOTrainer(...) trainer.add_callback(RichProgressCallback())需要注意的是部分回调如LogCompletionsCallback、WeaveCallback在构造时需要传入trainer实例因此只能先创建 Trainer 再通过add_callback挂载而BEMACallback、RichProgressCallback与训练器本身无依赖关系两种方式皆可。RichProgressCallback用 Rich 展示训练与评估实时进度RichProgressCallback是一个基于 Rich 库的进度展示回调它在终端中以实时刷新的方式呈现训练与评估进度是disable_tqdmTrue场景下的理想替代品。基本用法from trl import RichProgressCallback trainer DPOTrainer(...) trainer.add_callback(RichProgressCallback())该回调在构造时要求环境中已安装rich否则会抛出ImportError并提示pip install rich见 trl/trainer/callbacks.py。这与 transformers 的is_rich_available()检测逻辑保持一致属于可选依赖。源码层面的实现细节从 trl/trainer/callbacks.py 可以看到该回调利用 Rich 的Progress、Live、Panel、Table、Columns组件搭建了一个分组仪表盘其行为可以拆解为on_train_begin在is_world_process_zero仅主进程时创建训练进度条、评估进度条与状态面板训练进度条的总量取自state.max_steps并标注为蓝色[blue]Training。on_step_end按state.global_step与上次记录步数的差值推进训练进度条。on_prediction_step若评估 dataloader 可计算长度has_length则创建/推进评估进度条。on_log这是最有价值的部分。回调会将logs字典中形如train/loss、eval/accuracy的键按/前缀分组为每组生成一张独立的指标表格浮点数保留 3 位小数再用Columns并排布局最终包在标题为Step {global_step}的绿色面板中实时更新。on_train_end/on_evaluate/on_predict负责停止 Live 面板并清理内部状态。得益于这种按前缀分组的渲染逻辑DPO 的train/loss、GRPO 的objective/kl、eval/*等指标会自动分栏展示方便在终端中同时观察多个维度的指标走向。测试验证tests/test_rich_progress_callback.py 中通过一个DummyModel配eval_strategysteps、eval_steps1、logging_steps1、disable_tqdmTrue完成了一轮完整训练覆盖了训练与评估进度条、指标分组的完整生命周期可用作接入该回调的最小参考配置。LogCompletionsCallback将模型生成结果记录到 WB 与 CometLogCompletionsCallback会在训练过程中周期性驱动当前模型对评估集的 prompt 做生成并把 prompt completion 以表格形式记录到 Weights Biases 或 Comet用于定性观察模型输出质量的演化。这是 RL/偏好对齐训练中最实用的调试手段之一。构造参数签名定义于 trl/trainer/callbacks.py参数类型默认值说明trainerTrainer必填回调挂载的 Trainer其评估数据集必须包含prompt列否则构造时抛出ValueErrorgeneration_configGenerationConfigNone生成 completion 时使用的生成配置num_promptsintNone参与生成的 prompt 数量不传则使用整个评估集freqintNone每隔多少个 step 记录一次不传则跟随训练器的eval_steps使用示例from transformers import GenerationConfig from trl import LogCompletionsCallback trainer DPOTrainer(...) completions_callback LogCompletionsCallback( trainertrainer, generation_configGenerationConfig(max_length32), num_prompts8, # 只对评估集中前 8 条 prompt 生成 freq4, # 每 4 个 step 记录一次 ) trainer.add_callback(completions_callback)底层工作流程on_step_end中的处理逻辑trl/trainer/callbacks.py大致如下去重通过_last_logged_step保证同一 step 只记录一次该钩子可能被多次调用。频率控制freq self.freq or state.eval_steps仅当global_step % freq 0时触发记录。切分与生成使用accelerator.split_between_processes将评估集的prompt列切分到各进程经maybe_apply_chat_template套用聊天模板后调用模块级函数_generate_completions批量生成trl/trainer/callbacks.py。_generate_completions会通过unwrap_model_for_generation解包并行封装按per_device_eval_batch_size分 batch 调用model.generate再以去掉 prompt 部分、skip_special_tokensTrue解码的方式得到纯 completion。聚合gather_object汇总各进程的 prompts 与 completions。记录仅主进程构造pandas.DataFrame列为step/prompt/completion且跨记录步累积成历史表格根据report_to配置写入 wandbwandb.log({completions: table})或 Comet调用 trl/trainer/utils.py 的log_table_to_comet_experiment以completions.csv文件名落盘。注意源码中会将tokenizer.padding_side临时设为left这是为了确保 batch 内不同长度的 prompt 在左填充下完成正确对齐避免右侧填充干扰生成输出。测试验证tests/test_callbacks.py 提供了两个端到端测试test_basic_wandb验证 wandb 记录下的表格包含step/prompt/completion三列且prompt内容与评估集原始样本一致num_prompts2。test_basic_comet验证 Comet 实验的 asset 中按预期次数产出了completions.csv文件。这两个测试用require_wandb/require_comet装饰器控制依赖可作为理解回调行为的活文档。BEMACallback偏差校正指数移动平均BEMABEMACallback实现了 BEMABias-Corrected Exponential Moving Average算法即对训练权重维护一条偏差校正的指数移动平均轨迹周期性把running_model权重更新为该轨迹训练结束时将这条平滑权重保存为output_dir/bema下的完整模型用于推理或评估。它在 EMA 的基础上引入一个随步数衰减的缩放因子从而校正 EMA 的初始偏差。核心数学原理设当前模型权重为θ_tθ_0为第一个更新步update_after时的权重快照EMA_t为指数移动平均则 BEMA 权重为θ_t α_t · (θ_t - θ_0) EMA_t其中 EMA 按衰减因子β_t递推更新EMA_t (1 - β_t) · EMA_{t-1} β_t · θ_t两个时变因子都随步数t以幂律形式衰减α_t (ρ γ · t)^(-η) β_t (ρ γ · t)^(-κ)从实现上看trl/trainer/callbacks.py_ema_beta与_bema_alpha正是上述两条公式的直接翻译且β_t额外受min_ema_multiplier下限约束。构造参数参数默认值论文记号说明update_freq400φ每多少个 step 更新一次 BEMA 权重ema_power0.5κEMA 衰减因子的幂指数设为0.0可禁用 EMAbias_power0.2ηBEMA 缩放因子的幂指数大值如8.0令α_t迅速衰减到 0近似关闭偏置校正0.0则令α_t恒为 1最大、不衰减的校正lag10ρ权重衰减调度中的初始偏移充当更新的虚拟起始年龄控制早期平滑度update_after0τ预热步数在此之前不更新 BEMA 权重multiplier1.0γEMA 衰减因子的初始乘数min_ema_multiplier0.0-EMA 衰减因子的下限devicecpu-BEMA 缓冲区的承载设备。文档与源码都强调多数情况下该设备应与训练设备不同如训练在 CUDA、缓冲区放 CPU以避免显存 OOM生命周期行为on_train_begin通过_unwrap_model解包 DeepSpeed/FSDP/DataParallel 等封装用type(model)(model.config)克隆一份running_model并加载当前权重随后缓存所有requires_gradTrue的参数并在self.device上为每个参数克隆θ_0与EMA缓冲区EMA 初始化为θ_0。这正是参数表中缓冲区独立于训练设备的落地方式。on_step_end步数 update_after时跳过在step update_after时把θ_0与EMA快照为当前权重此后每满update_freq步执行_update_bema_weights在torch.no_grad()下完成 EMA 递推与run_param ema α·(θ_t - θ_0)的写入。on_train_end主进程将running_model保存到{output_dir}/bema并以日志提示保存路径。使用示例与测试验证from trl import BEMACallback trainer Trainer(..., callbacks[BEMACallback(update_freq400)])tests/test_callbacks.py 中的TestBEMACallback提供了多组验证test_model_saved训练结束后{output_dir}/bema目录存在且可用AutoModelForCausalLM.from_pretrained直接加载。test_update_frequency_0/1/2通过 mock_update_bema_weights精确断言触发步数。例如 9 个 step、update_freq2时在 step2, 4, 6, 8触发update_freq3时在3, 6, 9触发update_freq2, update_after3时从 step 3 之后按每 2 步触发5, 7, 9。test_bias_power_zero与test_no_ema分别验证bias_power0.0不衰减偏置校正与ema_power0.0禁用 EMA两种极端配置可正常运行。扩展将 BEMA 权重同步给参考模型trl/experimental/bema_for_ref_model/callback.py 中定义了BEMACallback的扩展版本额外提供update_ref_model、ref_model_update_freq、ref_model_update_after三个参数当update_ref_modelTrue时会在每次 BEMA 权重更新后把running_model的 state dict 写入参考模型ref_model通过扩展的CallbackHandlerWithRefModel以ref_modelkwarg 注入从而构造一个主模型的滞后平滑版本作为参考模型。它同时处理了 PEFTlora_/adapter_参数过滤、get_base_model解包与全量微调两种情形。这一能力在依赖参考模型提供 log-prob 的 DPO 类训练中具有实际价值。WeaveCallback面向 WB Weave 的追踪与评估日志WeaveCallback将每次评估阶段的 prompt 与模型生成结果写入 WB Weave并根据是否提供scorers分为两种模式见 trl/trainer/callbacks.py追踪模式Tracing ModescorersNone仅记录预测prompt → completion用于数据探索与分析。评估模式Evaluation Mode提供scorers在记录预测的同时用自定义评分函数对每条输出打分并汇总平均分等摘要指标。两种模式都基于 Weave 的EvaluationLogger实现结构化、一致的日志。与LogCompletionsCallback在on_step_end触发不同它只在评估阶段on_evaluate记录语义上更贴合按评估周期观察模型质量的需求也更高效。构造参数参数默认值说明trainer必填挂载的 Trainer其评估数据集必须含prompt列否则抛ValueErrorproject_nameNoneWeave 项目名缺省时依次尝试复用已有 weave client、当前 wandb run 的entity/project均不可用则抛错scorersNone{名字: 评分函数}字典评分函数签名为scorer(prompt: str, completion: str) - float提供则进入评估模式generation_configNone生成配置num_promptsNone参与评估的 prompt 数量上限dataset_nameeval_datasetWeave 中数据集元数据名称model_nameNoneWeave 中模型元数据名称缺省时从模型 config 的_name_or_path提取使用示例from trl import WeaveCallback # 追踪模式只记录预测 weave_callback WeaveCallback(trainertrainer, project_namemy-llm-training) trainer.add_callback(weave_callback) # 评估模式记录预测 评分 摘要 def accuracy_scorer(prompt: str, completion: str) - float: # 你的评分逻辑可通过 eval_attributes 访问元数据 return score weave_callback WeaveCallback( trainertrainer, project_namemy-llm-training, # 仅在 weave client 未初始化时需要 scorers{accuracy: accuracy_scorer}, ) trainer.add_callback(weave_callback)容错设计源码对缺失依赖与局部失败做了多层保护_initialize_weave在 weave 未安装时仅记录 warning 并跳过日志单条预测记录失败、单个 scorer 抛异常、摘要日志失败时都只记录 warning 而不中断训练on_evaluate也通过_last_logged_step去重避免同一 step 的重复记录。评估模式下eval_logger.log_summary会输出total_predictions、successful_predictions以及每个 scorer 的avg_{scorer_name}汇总统计便于在 Weave 仪表盘上横向对比各评估周期的得分趋势。SyncRefModelCallback训练过程中同步参考模型虽然docs/source/callbacks.md的 API 列表未显式列出SyncRefModelCallback但它是 trl/trainer/callbacks.py 中承担关键职责的基础回调被多个 Trainer 在内部自动挂载理解它对排查 DPO/GRPO 类训练行为很有帮助。工作机制构造参数为ref_model与accelerator。在on_step_end中当global_step % args.ref_model_sync_steps 0时对参考模型执行指数混合同步_sync_target_modeltarget_param.data.mul_(1.0 - alpha).add_(copy_param.data, alphaalpha)即以ref_model_mixup_alpha为混合系数把主模型权重按比例融入参考模型。其特殊之处在于对 DeepSpeed ZeRO-3 的处理若检测到zero_stage 3会通过deepspeed.zero.GatheredParameters收集主模型与参考模型的参数后在 rank 0 上执行更新保证分片环境下同步正确性。在训练器中的实际挂载该回调由训练器自动添加用户无需手动介入trl/trainer/dpo_trainer.pyDPOTrainer在sync_ref_modelTrue时挂载并校验其与 PEFT、precompute_ref_log_probsTrue的兼容性PEFT 下参考模型以禁用 adapter 的方式恢复不存在独立ref_model实例预计算参考 log-prob 假设参考模型固定二者均与周期性同步冲突。trl/trainer/grpo_trainer.pyGRPOTrainer以同样方式挂载。trl/trainer/kto_trainer.py 与 trl/trainer/rloo_trainer.pyKTO、RLOO 训练器同样使用。此外trl/experimental/sdpo/teacher_sync.py 与 trl/experimental/sdft/teacher_sync.py 分别将其子类化为SyncTeacherModelCallback用于在 SDPO/SDFT 实验中按配置的同步步数把学生模型同步为 EMA 教师模型。从源码结构看这一回调是 TRL 中参考模型动态更新机制的通用底座。回调组合使用与实战建议几个回调的触发时机互补可以自由组合RichProgressCallbackdisable_tqdmTrue无头/少依赖环境下获得更美观的终端面板适合日常本地调试。LogCompletionsCallback或WeaveCallback二选一两者都做评估集生成 外部平台记录前者面向 WB/Comet 表格后者面向 Weave 的 trace/评估体系若同时使用会重复生成按你的观测平台取舍即可。BEMACallback作为训练的无侵入附加物最终产物output_dir/bema独立于 checkpoint 体系可直接用于模型评测记得把device设为与训练设备不同的设备如 CPU以避免显存压力。SyncRefModelCallback由 DPO/GRPO/KTO/RLOO 训练器自动管理日常无需手动实例化。所有回调的完整实现、参数与生命周期钩子均可回溯至 trl/trainer/callbacks.py端到端行为可通过 tests/test_callbacks.py 与 tests/test_rich_progress_callback.py 中的测试用例复现在 DPO/GRPO/KTO/RLOO 各训练器源码中搜索add_callback可以进一步查看它们的自动挂载上下文。【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表