尧图精选

SuperGradients Phase Callbacks 完全指南:在训练管线任意节点注入自定义逻辑

🕒 发布时间:2026/9/18 23:12:16 📁 来源:尧图网络
SuperGradients Phase Callbacks 完全指南在训练管线任意节点注入自定义逻辑【免费下载链接】super-gradientsEasily train or fine-tune SOTA computer vision models with one open source training library. The home of Yolo-NAS.项目地址: https://gitcode.com/GitHub_Trending/su/super-gradientsPhase Callbacks阶段回调是 SuperGradients 训练框架提供的钩子hook机制用户可以将任意可调用对象注册到训练管线的特定时间点如训练开始、每个 batch 的 loss 计算之后、每个 epoch 结束、验证最佳模型出现时等从而实现学习率调度、数据增强切换、结果可视化、ONNX 导出校验等自定义行为。本文以 documentation/source/PhaseCallbacks.md 为主体结合仓库源码base_callbacks.py、callbacks.py 等深入讲解回调的底层事件模型、PhaseContext上下文对象、内置回调清单并通过一个完整的保存首个 batch 图像示例带你掌握在 Python 脚本与 YAML Recipe 两种方式下编写和使用自定义回调的完整流程。一、为什么需要 Phase Callbacks把自定义逻辑缝进训练管线将自有代码集成进一个现成的训练管线往往需要用户付出大量精力要么侵入式修改训练循环要么复制粘贴整个训练脚本再改动。SuperGradients 通过training_params.phase_callbacks参数解决了这一问题——你可以在调用Trainer.train(...)时传入一组在训练代码特定位置被触发的可调用对象callable训练循环本身保持不动所有定制逻辑都以回调的形式被外挂进去。在 sg_trainer.py 中这些回调会被收集进一个CallbackHandler统一驱动见下文回调如何被驱动一节而整个事件分发机制、Phase枚举与PhaseContext的实现都位于 base_callbacks.py。二、内置回调一览super_gradients.training.utils.callbacks模块为常见需求实现了开箱即用的回调完整导出清单见 callbacks/init.py回调类用途ModelConversionCheckCallback训练开始前将模型转为 ONNX 并用 ONNX Runtime 验证输出一致性提前暴露导出问题DeciLabUploadCallback训练结束后上传模型到 Deci Lab 并触发优化流程LRCallbackBase/LinearEpochLRWarmup/LinearBatchLRWarmup学习率 warmup 的基类与两种线性 warmup 实现按 epoch / 按 batch 步进StepLRScheduler/ExponentialLRScheduler/PolyLRScheduler/CosineLRScheduler/FunctionLRScheduler阶梯、指数、多项式、余弦退火与自定义函数五种硬编码 LR 调度策略LRSchedulerCallback包装任意torch.optim.lr_scheduler支持按metric_name驱动ReduceLROnPlateauDetectionVisualizationCallback/BinarySegmentationVisualizationCallback训练/验证过程中将检测框、二值分割预测结果可视化并写入sg_loggerTrainingStageSwitchCallbackBase/YoloXTrainingStageSwitchCallback多阶段训练切换基类YoloX 专用在指定 epoch 关闭增强并开启 L1 lossTimerCallback统计各阶段耗时并写入日志如train_batch_forward_with_loss_msSlidingWindowValidationCallback在最后一个 epoch 对验证集/测试集启用滑窗推理RoboflowResultCallback训练结束后把数据集名与 mAP 追加写入 CSVMetricsUpdateCallback/KDModelMetricsUpdateCallback在指定阶段更新指标计算器与 loss 平均表KD 版针对蒸馏场景使用学生输出PPYoloETrainingStageSwitchCallbackPP-YOLOE 的训练阶段切换同样导出自 callbacks/init.pyExtremeBatchPoseEstimationVisualizationCallback姿态估计任务中的批数据可视化2.1 典型案例YoloX 训练阶段切换YoloX 的 COCO 检测训练 Recipe 使用YoloXTrainingStageSwitchCallback在 epoch 285 起关闭数据增强并启用 L1 loss。完整配置见 coco2017_yolox_train_params.yamldefaults: - default_train_params max_epochs: 300 lr_mode: CosineLRScheduler cosine_final_lr_ratio: 0.05 lr_warmup_epochs: 5 lr_cooldown_epochs: 15 initial_lr: 0.02 zero_weight_decay_on_bias_and_bn: True batch_accumulate: 1 save_ckpt_epoch_list: [285] loss: YoloXDetectionLoss criterion_params: strides: [8, 16, 32] # output strides of all yolo outputs num_classes: 80 optimizer: SGD optimizer_params: momentum: 0.9 weight_decay: 0.0005 nesterov: True ema: True mixed_precision: True phase_callbacks: - YoloXTrainingStageSwitchCallback: next_stage_start_epoch: 285从源码看YoloXTrainingStageSwitchCallbackcallbacks.py继承自TrainingStageSwitchCallbackBase其apply_stage_change会遍历context.train_loader.dataset.transforms对具备close方法的 transform即 Mosaic 等需要显式关闭的增强逐一调用close()重置数据加载器迭代器并将context.criterion.use_l1 True切换为 L1 损失register_callback(Callbacks.YOLOX_TRAINING_STAGE_SWITCH) class YoloXTrainingStageSwitchCallback(TrainingStageSwitchCallbackBase): def __init__(self, next_stage_start_epoch: int 285): super(YoloXTrainingStageSwitchCallback, self).__init__(next_stage_start_epochnext_stage_start_epoch) def apply_stage_change(self, context: PhaseContext): for transform in context.train_loader.dataset.transforms: if hasattr(transform, close): transform.close() iter(context.train_loader) context.criterion.use_l1 True另一个常见用法是BinarySegmentationVisualizationCallback在训练期间可视化预测结果完整示例可见仓库中的分割迁移学习 Notebooktransfer_learning_semantic_segmentation.ipynb。三、回调如何工作完整事件模型所有回调都继承自 base_callbacks.py 中定义的Callback基类。基类为训练管线的每一个关键节点都声明了一个默认空实现的方法# super_gradients.training.utils.callbacks.base_callbacks.Callback class Callback: def on_training_start(self, context: PhaseContext) - None: pass def on_train_loader_start(self, context: PhaseContext) - None: pass def on_train_batch_start(self, context: PhaseContext) - None: pass def on_train_batch_loss_end(self, context: PhaseContext) - None: pass def on_train_batch_backward_end(self, context: PhaseContext) - None: pass def on_train_batch_gradient_step_start(self, context: PhaseContext) - None: pass def on_train_batch_gradient_step_end(self, context: PhaseContext) - None: pass def on_train_batch_end(self, context: PhaseContext) - None: pass def on_train_loader_end(self, context: PhaseContext) - None: pass def on_validation_loader_start(self, context: PhaseContext) - None: pass def on_validation_batch_start(self, context: PhaseContext) - None: pass def on_validation_batch_end(self, context: PhaseContext) - None: pass def on_validation_loader_end(self, context: PhaseContext) - None: pass def on_validation_end_best_epoch(self, context: PhaseContext) - None: pass def on_test_loader_start(self, context: PhaseContext) - None: pass def on_test_batch_start(self, context: PhaseContext) - None: pass def on_test_batch_end(self, context: PhaseContext) - None: pass def on_test_loader_end(self, context: PhaseContext) - None: pass def on_training_end(self, context: PhaseContext) - None: pass3.1 事件触发顺序各事件的调用顺序如下注释来自基类文档字符串见 base_callbacks.pyon_training_start(context) # 训练开始前调用一次适合设置 warmup LR for epoch in range(epochs): on_train_loader_start(context) for batch in train_loader: on_train_batch_start(context) on_train_batch_loss_end(context) # loss 计算完成后调用 on_train_batch_backward_end(context) # .backward() 调用之后 on_train_batch_gradient_step_start(context) # optimizer step 之前可用于梯度裁剪、梯度日志 on_train_batch_gradient_step_end(context) # 梯度更新之后适合基于 step 的调度器更新 LR on_train_batch_end(context) on_train_loader_end(context) on_validation_loader_start(context) for batch in validation_loader: on_validation_batch_start(context) on_validation_batch_end(context) on_validation_loader_end(context) on_validation_end_best_epoch(context) on_test_start(context) for batch in test_loader: on_test_batch_start(context) on_test_batch_end(context) on_test_end(context) on_training_end(context) # 训练结束后调用一次此外基类还定义了on_average_best_models_validation_start与on_average_best_models_validation_end两个事件用于平均最优模型EMA 类验证阶段的通知。自定义回调只需继承Callback并覆写其中需要的方法即可未被覆写的方法默认不做任何事。3.2 事件与 Phase 枚举的对应关系从源码结构看旧式PhaseCallback通过Phase枚举把事件浓缩为 13 个阶段base_callbacks.pyPRE_TRAINING - on_training_start TRAIN_EPOCH_START - on_train_loader_start TRAIN_BATCH_END - on_train_batch_loss_end TRAIN_BATCH_STEP - on_train_batch_gradient_step_end TRAIN_EPOCH_END - on_train_loader_end VALIDATION_BATCH_END - on_validation_batch_end VALIDATION_EPOCH_END - on_validation_loader_end VALIDATION_END_BEST_EPOCH - on_validation_end_best_epoch TEST_BATCH_END - on_test_batch_end TEST_END - on_test_loader_end AVERAGE_BEST_MODELS_VALIDATION_START - on_average_best_models_validation_start AVERAGE_BEST_MODELS_VALIDATION_END - on_average_best_models_validation_end POST_TRAINING - on_training_endPhaseCallback在 base_callbacks.py 中保留用于向后兼容旧代码它在构造函数中接收一个phase字符串或Phase枚举成员并在对应的on_*方法中转发调用自身的__call__(context)。新代码建议直接使用Callback基类。四、PhaseContext回调手中的训练现场快照你可能会注意到Callback的每个方法都只接收一个参数——PhaseContext实例。它代表了训练在某一时刻的完整状态快照包含大量训练属性见 base_callbacks.py 中PhaseContext.__init__的定义训练进度epoch、batch_idx模型与优化net模型、ema_modelEMA 模型、optimizer、criterion损失函数、lr_warmup_epochs数据inputs、preds、target、train_loader、valid_loader、test_loader指标与损失metrics_dict、metrics_compute_fn、loss_avg_meter、loss_log_items、loss_logging_items_names、valid_metrics、metric_to_watch实验与配置experiment_name、ckpt_dir、training_params、checkpoint_params、architecture、arch_params、device、ddp_silent_mode日志sg_logger控制与扩展stop_training置为True可请求提前停止训练、additional_batch_items额外 batch 数据每个属性默认均为None直到训练管线在对应节点将其计算或定义出来。例如on_training_start发生在第一个 epoch 开始之前因此此时context.epoch为None而optimizer、criterion、device、experiment_name、ckpt_dir、net、sg_logger、train_loader、valid_loader、training_params、checkpoint_params、arch_params、metric_to_watch、valid_metrics等则已经就绪。每个方法的 docstring 都会列出该时间点可用的上下文属性编写回调时请以对应方法的 docstring 为准。例如on_training_start的签名如下base_callbacks.pydef on_training_start(self, context: PhaseContext) - None: Called once before start of the first epoch At this point, the context argument will have the following attributes: - optimizer - criterion - device - experiment_name - ckpt_dir - net - sg_logger - train_loader - valid_loader - training_params - checkpoint_params - arch_params - metric_to_watch - valid_metrics The corresponding Phase enum value for this event is Phase.PRE_TRAINING. :param context: passPhaseContext还提供update_context(**kwargs)方法训练循环在推进时会不断用它刷新上下文内容回调中也可以主动调用它来更新/注入自己的数据例如让后面的回调读取。4.1 回调如何被驱动CallbackHandler从源码结构看训练循环并不直接逐个调用回调而是通过CallbackHandlerbase_callbacks.py统一驱动CallbackHandler本身也继承Callback在其构造函数中接收回调列表并把每一个on_*事件按列表顺序依次转发给所有回调class CallbackHandler(Callback): def __init__(self, callbacks: List[Callback]): self.callbacks callbacks def on_training_start(self, context: PhaseContext) - None: for callback in self.callbacks: callback.on_training_start(context) # ... 每个 on_* 事件均如此循环转发源码注释中还提到未来会引入优先级排序机制Forward/Loss/Backward/Metrics/Scheduler/Logging以保证多个相互依赖的回调如先梯度裁剪再记录梯度按正确顺序执行当前版本则严格按用户传入phase_callbacks列表的顺序调用。五、实战编写第一个自定义回调下面实现一个简单回调在每个 epoch 将训练与验证的首个 batch 图像保存到本地 checkpoints 目录下新建的batch_images文件夹中。它需要被触发 3 次训练开始时在 checkpoints 目录下创建batch_images文件夹每个训练 batch 进入网络之前若为 epoch 首个 batch 则保存图像每个验证 batch 进入网络之前保存首个 batch 图像。因此需要覆写Callback的on_training_start、on_train_batch_start和on_validation_batch_start三个方法from super_gradients.training.utils.callbacks import Callback, PhaseContext from super_gradients.common.environment.ddp_utils import multi_process_safe import os from torchvision.utils import save_image class SaveFirstBatchCallback(Callback): def __init__(self): self.outputs_path None self.saved_first_validation_batch False multi_process_safe def on_training_start(self, context: PhaseContext) - None: outputs_path os.path.join(context.ckpt_dir, batch_images) os.makedirs(outputs_path, exist_okTrue) multi_process_safe def on_train_batch_start(self, context: PhaseContext) - None: if context.batch_idx 0: save_image(context.inputs, os.path.join(self.outputs_path, ffirst_train_batch_epoch_{context.epoch}.png)) multi_process_safe def on_validation_batch_start(self, context: PhaseContext) - None: if context.batch_idx 0 and not self.saved_first_validation_batch: save_image(context.inputs, os.path.join(self.outputs_path, ffirst_validation_batch_epoch_{context.epoch}.png)) self.saved_first_validation_batch True5.1 重要DDP 多节点下的multi_process_safe当使用多节点训练时参见 device.md 中关于 DDP 的说明回调会在每个节点上各触发一次。这在某些场景下可能有意义但通常你希望每个 step 只触发一次。给方法加上multi_process_safe装饰器可以保证只在主进程rank 0触发。从源码看multi_process_safe实现在 ddp_utils.py它检查device_config.assigned_rank 0非主进程直接跳过函数体适用于无返回值的函数def multi_process_safe(func): def do_nothing(*args, **kwargs): pass wraps(func) def wrapper(*args, **kwargs): if device_config.assigned_rank 0: return func(*args, **kwargs) else: return do_nothing(*args, **kwargs) return wrapper在上面的例子中我们希望每个 step 只触发一次因此三个方法都加了该装饰器。六、在 Python 脚本中使用自定义回调自定义回调可以直接通过training_params.phase_callbacks传入Trainer.train(...)trainer Trainer(my_experiment) train_dataloader ... valid_dataloader ... model ... train_params { loss: CrossEntropyLoss, criterion_params: {}, phase_callbacks: [SaveFirstBatchCallback()], ... } trainer.train(training_paramstrain_params, train_loadertrain_dataloader, valid_loadervalid_dataloader)这种方式最直接回调对象是代码中的实例不涉及任何反序列化适合脚本化、实验性的训练流程。七、在 RecipeYAML 配置中使用自定义回调如果你使用配置文件工作详见 configuration_files.md则多一个注册步骤。这与在 Recipe 中使用任何自定义对象类似需要用register_callback装饰器把新回调注册进 SuperGradients 的注册表registry框架才能从.yamlRecipe 中实例化它。7.1 用register_callback注册from super_gradients.training.utils.callbacks import Callback, PhaseContext from super_gradients.common.environment.ddp_utils import multi_process_safe import os from torchvision.utils import save_image from super_gradients.common.registry.registry import register_callback register_callback() class SaveFirstBatchCallback(Callback): def __init__(self): self.outputs_path None self.saved_first_validation_batch False multi_process_safe def on_training_start(self, context: PhaseContext) - None: outputs_path os.path.join(context.ckpt_dir, batch_images) os.makedirs(outputs_path, exist_okTrue) multi_process_safe def on_train_batch_start(self, context: PhaseContext) - None: if context.batch_idx 0: save_image(context.inputs, os.path.join(self.outputs_path, ffirst_train_batch_epoch_{context.epoch}.png)) multi_process_safe def on_validation_batch_start(self, context: PhaseContext) - None: if context.batch_idx 0 and not self.saved_first_validation_batch: save_image(context.inputs, os.path.join(self.outputs_path, ffirst_validation_batch_epoch_{context.epoch}.png)) self.saved_first_validation_batch True注册表定义见 registry.py框架内置回调如YoloXTrainingStageSwitchCallback、CosineLRScheduler等也正是通过register_callback/register_lr_scheduler/register_lr_warmup这些装饰器注册的注册后的名称会进入CALLBACKS等对象名注册表见 object_names.py。7.2 在 YAML 中声明回调然后在你的my_training_hyperparams.yaml中像使用任何 SG 内置 phase callback 一样声明它defaults: - default_train_params max_epochs: 250 ... phase_callbacks: - SaveFirstBatchCallback如果需要传参使用键值对形式参考 YoloX 示例的写法phase_callbacks: - SaveFirstBatchCallback: some_param: value7.3 在启动脚本中导入回调类最后务必在启动训练的脚本中 importSaveFirstBatchCallback。这是必须的——否则该类根本不会被加载SuperGradients 也就无法识别并实例化它这与在 Recipe 中使用自定义数据集、自定义损失时的要求一致from omegaconf import DictConfig import hydra import pkg_resources from my_callbacks import SaveFirstBatchCallback from super_gradients import Trainer, init_trainer hydra.main(config_pathpkg_resources.resource_filename(super_gradients.recipes, ), version_base1.2) def main(cfg: DictConfig) - None: Trainer.train_from_config(cfg) def run(): init_trainer() main() if __name__ __main__: run()八、深入理解内置回调族源码剖析为帮助你在自定义时找准参考实现这里剖析三类典型内置回调的实现要点8.1 LR 调度回调族LRCallbackBasecallbacks.py是所有硬编码 LR 调度器的基类它把initial_lr规范化为按参数组param group名称索引的字典并提供update_lr把计算好的self.lr写回optimizer.param_groups中每个分组。子类需实现两个抽象方法is_lr_scheduling_enabled(context)判断当前是否需要调度例如 warmup 尚未结束则跳过perform_scheduling(context)根据context.epoch/context.batch_idx计算新的 LR。以CosineLRScheduler为例callbacks.py它在TRAIN_BATCH_STEP阶段触发先扣除 warmup 与 cooldown 的影响计算有效 epoch/iter再用compute_learning_rate按余弦曲线从initial_lr衰减到initial_lr * cosine_final_lr_ratio。StepLRScheduler则在TRAIN_EPOCH_END触发按lr_updates里程碑逐次乘以lr_decay_factor。这些回调的名字如CosineLRScheduler正是 Recipe 中lr_mode: CosineLRScheduler字段所引用的对象。8.2 可视化回调族DetectionVisualizationCallback与BinarySegmentationVisualizationCallbackcallbacks.py都接收phase、freq每多少 epoch 触发一次、batch_idx对第几个 batch 可视化、last_img_idx_in_batch记录到第几张图默认 -1 记录整个 batch等参数。触发条件为context.epoch % self.freq 0 and context.batch_idx self.batch_idx随后把预测结果经post_prediction_callback检测场景后画到context.inputs上最终通过context.sg_logger.add_images(...)写入日志。8.3 阶段切换回调族TrainingStageSwitchCallbackBasecallbacks.py在TRAIN_EPOCH_START阶段触发当context.epoch self.next_stage_start_epoch时调用子类实现的apply_stage_change(context)。YoloXTrainingStageSwitchCallback和PPYoloETrainingStageSwitchCallback分别按各自算法需求改写context中的 transform 与 criterion实现多阶段训练。九、小结编写回调的检查清单选对基类新代码继承Callback并覆写on_*方法旧式PhaseCallback仅用于向后兼容。确认上下文可用属性查看对应方法的 docstring确认你需要的PhaseContext属性在该事件节点已就绪例如epoch在on_training_start中为None。控制触发频率与条件利用context.batch_idx、context.epoch过滤触发时机避免每个 batch 都执行昂贵操作。DDP 安全多节点训练时对只需执行一次的方法加multi_process_safe。Recipe 集成三步register_callback()注册 → YAML 中按名声明 → 启动脚本中 import 该类。通过本文介绍的事件模型、PhaseContext与注册机制你可以像使用框架内置回调一样把学习率策略、数据增强切换、训练可视化乃至模型导出校验等任意逻辑以非侵入的方式接入 SuperGradients 训练管线。更深入的配置体系可继续阅读 configuration_files.md 与 device.md。【免费下载链接】super-gradientsEasily train or fine-tune SOTA computer vision models with one open source training library. The home of Yolo-NAS.项目地址: https://gitcode.com/GitHub_Trending/su/super-gradients创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联 返回资讯列表 →