资讯动态

自定义 Hugging Face Trainer:通过子类化改变训练行为

发布时间:2026/9/10 2:06:09 来源:尧图企业网站定制
自定义 Hugging Face Trainer通过子类化改变训练行为【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers训练循环是深度学习中最复杂、最难以调试的代码之一。Hugging Facetransformers库的Trainer把整个训练流程封装成一套开箱即用的 API但它并不要求你接受要么全用、要么全不用通过子类化Subclassing你可以在不改写整个训练循环的前提下精准修改数据加载、损失计算等关键环节的行为。本文基于 docs/source/en/trainer_customize.md 展开结合仓库源码系统讲解Trainer子类化的核心方法、适用场景与实现原理读完即可把这一模式应用到自己的训练代码中。子类化 vs. 回调先想清楚要改什么还是何时在动手覆写方法之前首先要明确你的需求属于哪一类子类化Subclassing修改的是训练循环的内容——即Trainer计算什么前向传播、损失计算、数据加载、优化步骤。典型例子包括自定义损失函数、按多步批量生成数据、跳过标签参与推理等。回调Callback控制的是时机与条件——即Trainer何时、是否执行某个动作日志记录、评估、早停。例如每个训练步结束时记录指标用回调即可无需触碰训练循环本身。对应到代码层面如果你想改变训练循环内部的行为就继承Trainer并覆写目标方法如果只是想在外围挂钩子请参考 Callbacks 指南 使用回调。选择正确的机制能让你的代码更简洁、更易维护。[!NOTE] 可覆写方法的完整清单以Trainer的 API 文档为准。私有方法以_开头如_save_checkpoint、_evaluate也可以覆写但它们属于内部实现细节可能在版本升级时无预警地改变签名或行为覆写前需谨慎评估维护成本。核心一覆写get_train_dataloader改变数据加载默认行为单步取批、用后即弃Trainer的标准训练数据加载流程是加载一个 batch → 训练它 → 丢弃 → 加载下一个 batch。默认实现位于 trainer.pydef get_train_dataloader(self) - DataLoader: Returns the training [~torch.utils.data.DataLoader]. Will use no sampler if train_dataset does not implement __len__, a random sampler (adapted to distributed training if necessary) otherwise. Subclass and override this method if you want to inject some custom behavior. if self.train_dataset is None: raise ValueError(Trainer: training requires a train_dataset.) return self._get_dataloader( datasetself.train_dataset, descriptionTraining, batch_sizeself._train_batch_size, sampler_fnself._get_train_sampler, is_trainingTrue, )该方法最终委托给内部辅助方法_get_dataloadertrainer.py后者负责组装torch.utils.data.DataLoader设置batch_size、collate_fn即self.data_collator、num_workers来自args.dataloader_num_workers默认 0、pin_memory、persistent_workers、prefetch_factor等参数并调用self.accelerator.prepare完成设备与分布式适配。训练循环在_inner_training_loop中只调用一次get_train_dataloader()trainer.py来构建 dataloader随后按步迭代取批。这就是加载一个 batch → 训练 → 丢弃 → 再加载这一行为模型。实战场景GRPO 的跨步批量生成什么时候需要覆写这个方法GRPO一种在线强化学习算法是典型例子它要先生成完整回复completions再在生成结果上训练。生成是自回归过程代价极高——一段 512 token 的 completion 大约需要 512 次串行前向传播而一个训练步只有一次前向传播。如果每个训练步都重新生成成本不可接受。trl.GRPOTrainer的做法是覆写get_train_dataloader把多个训练步的生成 prompt 一次性批量加载通过steps_per_generation参数把 batch size 放大若干倍。例如train_batch_size4、steps_per_generation8时dataloader 产出 batch size 为 32 的批次一次生成服务 8 个训练步生成成本降低 8 倍def get_train_dataloader(self): dataloader_params { batch_size: self._train_batch_size * self.args.steps_per_generation, # 唯一改动 ... }这个模式的核心思想是数据加载策略与训练步频解耦。默认 Trainer 的_train_batch_size由per_device_train_batch_sizetraining_args.py 中默认 8等参数推导而来覆写后你可以自由调整产出的 batch 结构只要保证后续训练逻辑能正确消费。覆写时的注意事项保持签名一致get_train_dataloader(self) - DataLoader返回一个合法的DataLoader_get_dataloader的结果或你自己构造的均可。若你的数据集不实现__len__如IterableDataset默认不会附加 sampler取样策略由_get_train_samplertrainer.py决定支持random默认、sequential、group_by_length、batch_rebalance等。分布式场景下accelerator.prepare会负责 batch 的切分与广播覆写时若绕过_get_dataloader需要自行处理这些细节。核心二覆写compute_loss改变损失计算默认行为直接取模型的交叉熵损失Trainer默认从模型输出中取回损失。绝大多数模型在forward时传入labels即可内部计算交叉熵损失compute_loss只是把它取出来。默认实现位于 trainer.pydef compute_loss( self, model: nn.Module, inputs: dict[str, torch.Tensor | Any], return_outputs: bool False, num_items_in_batch: torch.Tensor | int | None None, ) - torch.Tensor | tuple[torch.Tensor, Any]: How the loss is computed by Trainer. By default, all models return the loss in the first element. ... Subclass and override for custom behavior. If you are not using num_items_in_batch when computing your loss, make sure to overwrite self.model_accepts_loss_kwargs to False. Otherwise, the loss calculation might be slightly inaccurate when performing gradient accumulation. ... outputs model(**inputs) ... loss outputs[loss] if isinstance(outputs, dict) else outputs[0] return (loss, outputs) if return_outputs else loss默认实现会依次处理label smoothing通过label_smoother、用户自定义损失函数compute_loss_func若有、以及直接从模型输出取loss的兜底逻辑。为什么 DPO 必须覆写compute_lossDPODirect Preference Optimization的损失计算与标准交叉熵完全不同它度量的是策略模型policy model对被选中回复相对被拒绝回复的偏好强度且必须对照一个冻结的参考模型reference model模型不接收 labels只返回 logits由 DPO 自行计算 log-probschosen选中与 rejected被拒绝的回复会拼接在同一 batch 中参考模型单独计算自己的 log-probs最终损失是π_chosen、π_rejected、π_ref_chosen、π_ref_rejected四个量的函数。这些需求与默认的compute_loss完全不兼容因此trl.DPOTrainer覆写了该方法def compute_loss( self, model: PreTrainedModel | nn.Module, inputs: dict[str, torch.Tensor | Any], return_outputsFalse, num_items_in_batchNone, ) - torch.Tensor | tuple[torch.Tensor, dict[str, float]]: ... outputs model(**inputs) logits outputs.logits logps get_logps(logits, inputs) chosen_logps, rejected_logps logps.chunk(2, dim0) # batch 按 [chosen, rejected] 排列 ref_logits self.ref_model(**inputs).logits ref_logps get_logps(ref_logits, inputs) ref_chosen_logps, ref_rejected_logps ref_logps.chunk(2, dim0) # batch 按 [chosen, rejected] 排列 chosen_scores chosen_logps - ref_chosen_logps rejected_scores rejected_logps - ref_rejected_logps per_sequence_loss -F.logsigmoid(self.beta * chosen_scores - rejected_scores) loss per_sequence_loss.mean() return (loss, outputs) if return_outputs else loss注意几个关键细节返回值契约return_outputsTrue时返回(loss, outputs)元组否则只返回loss。training_steptrainer.py调用compute_loss后直接用返回的 loss 做反向传播因此覆写时务必保持这个契约。梯度累积的配合compute_loss的签名包含num_items_in_batch参数。若你的自定义损失不使用它务必把self.model_accepts_loss_kwargs置为False否则梯度累积时 loss 归一化可能不准确training_step中会依据该标志决定是否按gradient_accumulation_steps缩放 loss。模型本身不计算 loss覆写场景下模型往往只输出 logits 而不接受 labels你需要完全掌控从 logits 到标量 loss 的整条链路。完整的子类化骨架把两个核心方法组合起来一个自定义 Trainer 的骨架如下from transformers import Trainer class MyCustomTrainer(Trainer): def get_train_dataloader(self): # 自定义跨多个训练步批量取数降低生成/预处理成本 return super().get_train_dataloader() def compute_loss(self, model, inputs, return_outputsFalse, num_items_in_batchNone): # 自定义非标准损失如 DPO、GRPO 的偏好/奖励损失 outputs model(**inputs) loss my_custom_loss(outputs, inputs) return (loss, outputs) if return_outputs else loss实例化时传入TrainingArgumentsbatch size、gradient_accumulation_steps、dataloader_num_workers等均通过 training_args.py 配置即可直接使用training_args TrainingArguments( output_dir./output, per_device_train_batch_size4, gradient_accumulation_steps8, ) trainer MyCustomTrainer( modelmodel, argstraining_args, train_datasetdataset, ) trainer.train()延伸阅读更多真实世界的子类化示例查看 TRL 中~trl.GRPOTrainer与~trl.DPOTrainer如何扩展Trainer或参考 Axolotl 在src/axolotl/core/trainers中基于Trainer构建自定义训练器的做法。如果只需要在某个训练事件如训练步结束时记录指标发生时挂钩子请阅读 Callbacks 指南——它比子类化更轻量是时机控制的首选方案。完整的可覆写方法清单与签名见Trainer的 API 文档与源码中的 docstring源码中大量方法都标注了 Subclass and override this method if you want to inject some custom behavior。【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价