ARTICLE DETAIL

资讯详情

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

Lightning Fabric Callbacks 详解:为训练循环注入可插拔行为的轻量级回调系统

Lightning Fabric Callbacks 详解:为训练循环注入可插拔行为的轻量级回调系统 Lightning Fabric Callbacks 详解为训练循环注入可插拔行为的轻量级回调系统【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000 GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning本文以 Lightning Fabric 官方文档 Callbacks 章节 为核心结合本仓库中 Fabric 源码src/lightning/fabric/fabric.py、官方自建 Trainer 示例examples/fabric/build_your_own_trainer/trainer.py与单元测试tests/tests_fabric/test_fabric.py系统讲解如何在自定义训练循环中注册与调用回调以及如何基于回调机制构建属于你自己的 Trainer。读完本文你将掌握Fabric(callbacks...)与fabric.call(...)的完整用法、多回调协同原理与参数过滤机制并能在自己的训练脚本中落地实现。什么是 Fabric CallbackCallback回调是一种设计模式允许你或使用你代码的用户在不修改训练循环源码的前提下向训练循环中注入新的行为。在 Lightning Fabric 中回调就是任意一个普通的 Python 对象其上的方法名如on_train_batch_end充当钩子hook——当训练循环运行到特定位置时由调用方主动触发这些方法。这一设计与 Lightning Trainer 中基于继承Callback基类的回调不同Fabric 的回调不要求继承任何基类、不要求实现任何固定接口甚至不要求对象上的所有方法都带统一签名。它只约定一个核心交互协议——你通过fabric.call(name, ...)触发同名方法其余一切交给你的想象力。正如文档所强调的回调方法的代码与训练器代码完全解耦decoupled这正是它具备高度扩展性的根源。从源码看Fabric.__init__接受的callbacks参数可以是一个对象或对象列表其类型标注为Optional[Union[list[Any], Any]]见 src/lightning/fabric/fabric.py#L143内部由_configure_callbacks统一规范为列表存储src/lightning/fabric/fabric.py#L1256-L1260staticmethod def _configure_callbacks(callbacks: Optional[Union[list[Any], Any]]) - list[Any]: callbacks callbacks if callbacks is not None else [] callbacks callbacks if isinstance(callbacks, list) else [callbacks] callbacks.extend(_load_external_callbacks(lightning.fabric.callbacks_factory)) return callbacks也就是说传单个对象会被自动包装成列表未传则为空列表此外它还会通过 Python 的 entry points 机制lightning.fabric.callbacks_factory组加载第三方包注册的外部回调实现见 src/lightning/fabric/utilities/registry.py#L25-L50这让回调生态可以独立于主程序扩展。为你的训练循环添加回调接口文档给出的第一个场景是假设你希望任何人都能在一次训练迭代结束时运行任意代码。在 Fabric 中实现分为两步——定义回调对象、在循环中触发它。第一步定义回调回调的代码可以放在任何地方完全独立于训练循环# my_callbacks.py class MyCallback: def on_train_batch_end(self, loss, output): # 在训练步结束时运行任意代码 ...注意这里的loss与output只是方法签名真正的值由调用方通过fabric.call(...)传入回调本身不关心它们来自哪里。第二步注册回调并在循环中触发from lightning.fabric import Fabric # 回调代码可以生活在任何地方远离训练循环 from my_callbacks import MyCallback # 注册一个或多个回调 fabric Fabric(callbacks[MyCallback()]) # ... 模型、优化器与 DataLoader 的 setup 略 ... for iteration, batch in enumerate(train_dataloader): ... fabric.backward(loss) optimizer.step() # 在合适的位置让回调做任意处理 # 通过关键字参数把变量传给回调 fabric.call(on_train_batch_end, lossloss, output...)整个流程的核心是 Fabric.call 方法def call(self, hook_name: str, *args: Any, **kwargs: Any) - None: for callback in self._callbacks: method getattr(callback, hook_name, None) if method is None: continue if not callable(method): rank_zero_warn( fSkipping the callback {type(callback).__name__}.{hook_name} because it is not callable. ) continue filtered_kwargs self._filter_kwargs_for_callback(method, kwargs) method(*args, **filtered_kwargs)这段实现揭示了几个关键行为按注册顺序依次调用fabric.call遍历self._callbacks列表按传入Fabric(callbacks...)的顺序逐个触发缺失方法自动跳过某个回调没有名为hook_name的方法时getattr返回None直接跳过不会报错——因此并非所有注册对象都必须实现同名方法不可调用属性会告警并跳过若找到的属性存在但不是可调用对象callable会通过rank_zero_warn发出警告src/lightning/fabric/utilities/rank_zero.py 中定义仅在 rank 0 打印避免静默失败支持位置参数与关键字参数*args原样透传**kwargs则会经过签名过滤后再传入。练习原文档原题实现一个回调计算并打印完成一次迭代iteration所花费的时间。你可以利用on_train_batch_start记录起始时间、on_train_batch_end计算耗时这正好验证了回调接口的任意命名、任意时机触发特性。同时运行多个回调回调系统的设计目标之一就是让多个回调轻松并行工作。只需将多个实现放进一个列表传给Fabric# 以列表形式传入多个回调实现 callback1 LearningRateMonitor() callback2 Profiler() fabric Fabric(callbacks[callback1, callback2]) # 让 Fabric 调用各实现中如果存在的同名方法 fabric.call(any_callback_method, arg1..., arg2...) # 上面的 fabric.call 等价于手动依次调用 callback1.any_callback_method(arg1..., arg2...) callback2.any_callback_method(arg1..., arg2...)这里有几个要点fabric.call按回调被赋予 Fabric 时的顺序调用它们这一点由call方法的 for 循环顺序直接决定未实现同名方法的回调会被自动跳过因此你可以在列表中混入能力各异的回调互不干扰注意LearningRateMonitor、Profiler这类命名只是示意——Fabric 本身并不强制这些类名文档在此处用于说明不同回调可以各司其职。不同签名如何协同按签名自动过滤关键字参数多个回调可以拥有不同的方法签名。Fabric 会根据每个回调方法的函数签名自动过滤关键字参数从而让签名各异的回调无缝协同。这是 Fabric 回调系统最具特色的一点文档用三个回调给出了完整示例class TrainingMetricsCallback: def on_train_epoch_end(self, train_loss): print(fTraining loss: {train_loss:.4f}) class ValidationMetricsCallback: def on_train_epoch_end(self, val_accuracy): print(fValidation accuracy: {val_accuracy:.4f}) class ComprehensiveCallback: def on_train_epoch_end(self, epoch, **kwargs): print(fEpoch {epoch} complete with metrics: {kwargs}) fabric Fabric( callbacks[TrainingMetricsCallback(), ValidationMetricsCallback(), ComprehensiveCallback()] ) # 每个回调只会收到它能处理的参数 fabric.call(on_train_epoch_end, epoch5, train_loss0.1, val_accuracy0.95, learning_rate0.001)调用后三个回调分别获得TrainingMetricsCallback只收到train_loss0.1ValidationMetricsCallback只收到val_accuracy0.95ComprehensiveCallback收到epoch5并通过**kwargs捕获其余全部参数train_loss、val_accuracy、learning_rate。底层实现_filter_kwargs_for_callback这一智能分发能力的实现位于 src/lightning/fabric/fabric.py#L1012-L1038def _filter_kwargs_for_callback(self, method: Callable, kwargs: dict[str, Any]) - dict[str, Any]: try: sig inspect.signature(method) except (ValueError, TypeError): # 无法检查签名时为保持向后兼容透传所有 kwargs return kwargs filtered_kwargs {} for name, param in sig.parameters.items(): # 如果方法接受 **kwargs直接透传所有原始 kwargs if param.kind inspect.Parameter.VAR_KEYWORD: return kwargs # 如果参数名出现在传入的 kwargs 中则加入过滤结果 if name in kwargs: filtered_kwargs[name] kwargs[name] return filtered_kwargs其过滤规则可以总结为三点按参数名匹配回调方法签名中出现的参数名若在调用方传入的 kwargs 中存在则被保留未出现在签名中的 kwargs 被丢弃**kwargs通配只要方法声明了**kwargsinspect.Parameter.VAR_KEYWORD所有原始 kwargs 原样透传不做任何过滤签名检查失败则宽容处理当inspect.signature抛ValueError/TypeError例如某些动态生成或 mock 对象时为保持向后兼容直接返回全部 kwargs。注意过滤只针对关键字参数kwargs位置参数*args在fabric.call中是不加过滤、原样透传给每个回调的因此设计回调时建议优先使用关键字参数来接收数据。对应地仓库的单元测试 tests/tests_fabric/test_fabric.py#L1299-L1325test_callback_kwargs_filtering精确验证了这一行为CallbackWithLimitedKwargs只拿到epochCallbackWithVarKeywords拿到epoch和全部剩余 kwargsCallbackWithNoParams无参数也能被正常调用。同文件的test_callback_kwargs_filtering_signature_inspection_failureL1328-L1349则验证了签名检查失败时透传全部 kwargs 的兜底逻辑。这些测试用例本身即是理解回调协议行为的最佳参考。其他值得注意的回调触发点虽然fabric.call是你手动触发回调的入口但 Fabric 内部的若干机制也会自动调用回调理解这些能帮助你更完整地利用回调系统on_after_setup在fabric.setup(...)完成模型/优化器设置后被自动调用传入fabricself, modulemodule见 src/lightning/fabric/fabric.py#L308on_after_optimizer_step_FabricOptimizer.step()在每个优化器步之后自动遍历回调并触发见 src/lightning/fabric/wrappers.py#L91-L95传入strategy与optimizerLightningModule自动注册为回调当传入setup的对象具有_fabric属性即 LightningModule时它会被自动追加进self._callbacks列表src/lightning/fabric/fabric.py#L302-L306这意味着你的 LightningModule 上的同名 hook 方法也会被fabric.call一并触发——回调与模型 hook 共享同一套调用协议。实战用回调构建你自己的 Trainer文档在Next steps中指出回调是构建 Trainer 的强大工具并推荐参考官方基于 Fabric 的 Trainer 模板。该模板即本仓库的 examples/fabric/build_your_own_trainer/trainer.py配套运行入口见 examples/fabric/build_your_own_trainer/run.py其MyCustomTrainer.__init__直接透传callbacks给L.Fabricself.fabric L.Fabric( acceleratoraccelerator, strategystrategy, devicesdevices, precisionprecision, pluginsplugins, callbackscallbacks, loggersloggers, )随后在训练循环的各个关键位置调用fabric.call(...)见 trainer.py#L215-L260 的train_loop与 L262-L322 的val_loop支持的钩子包括on_train_epoch_start/on_train_epoch_endon_train_batch_start/on_train_batch_endon_before_backward/on_after_backwardon_before_zero_grad/on_before_optimizer_stepon_validation_model_eval/on_validation_model_trainon_validation_epoch_start/on_validation_epoch_endon_validation_batch_start/on_validation_batch_end例如在train_loop中self.fabric.call(on_train_epoch_start) for batch_idx, batch in enumerate(iterable): ... self.fabric.call(on_train_batch_start, batch, batch_idx) ... self.fabric.call(on_train_batch_end, self._current_train_return, batch, batch_idx) ... self.fabric.call(on_train_epoch_end)而在training_step中trainer.py#L324-L345围绕fabric.backward前后分别触发on_before_backward与on_after_backwardoutputs model.training_step(batch, batch_idxbatch_idx) loss outputs if isinstance(outputs, torch.Tensor) else outputs[loss] self.fabric.call(on_before_backward, loss) self.fabric.backward(loss) self.fabric.call(on_after_backward)这套模式展示了回调的典型落地方式训练循环只负责在正确的位置调用fabric.call具体做什么完全交给外部注册的回调——这正是文档所述让使用者无需修改源码即可扩展训练循环的工程实践。注意该模板的 docstring 特别警告——为 Lightning Trainer 编写的回调尤其依赖 trainer 对象的回调在 Fabric 下不会工作因为 Fabric 的回调协议中并没有一个统一的trainer全局对象回调只能拿到调用方显式传入的数据。这是 Fabric 回调与 LightningCallback基类之间最本质的差异。使用注意事项与最佳实践综合原文档与源码实现使用 Fabric 回调时有几点值得牢记回调就是普通对象无需继承、无需注册装饰器任何定义了钩子方法的对象都能成为回调方法缺失时fabric.call自动跳过因此可以为不同回调定义不同的方法集合。顺序即调用顺序回调按传入列表的顺序被调用若回调之间存在依赖如一个回调的产出被另一个消费请自行保证列表顺序。善用签名过滤给多个回调传同一批 kwargs 是安全且推荐的——每个回调只会收到自己签名需要的参数需要兜底接收一切时请声明**kwargs。若回调通过fabric.call(hook, arg)以位置参数方式接收数据则该参数会原样传给所有实现了该 hook 的回调请注意签名匹配。保持循环与回调解耦训练循环中只写fabric.call(...)具体逻辑放回调里这样你的代码使用者可以在不接触训练循环源码的前提下自由扩展行为——这正是回调机制的核心价值。用测试锚定行为本仓库的 tests/tests_fabric/test_fabric.py 是理解回调协议行为边界的绝佳参考多回调顺序、缺失方法跳过、kwargs 过滤、签名检查失败兜底等均有覆盖在实现自定义回调前不妨先通读一遍。如果你希望构建一个功能完整的 Trainer可以以 examples/fabric/build_your_own_trainer/trainer.py 为起点将它定制成适合自己需求的版本——回调机制正是这个模板保持小而可读的关键设计。【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000 GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表