完全指南:Recipe 配置到 Python 对象的自动实例化机制)
SuperGradients 配置工厂Factories完全指南Recipe 配置到 Python 对象的自动实例化机制【免费下载链接】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-gradientsSuperGradients 的 Factories工厂机制是连接 YAML 训练配方Recipe与 Python 对象实例之间的桥梁它让配置文件中以字符串/字典形式声明的模型、数据集、变换、损失函数等对象被自动识别并按需实例化从而无需编写大量样板代码。本指南以documentation/source/Recipes_Factories.md为主线系统讲解内置工厂的使用、自定义类的注册方法并结合src/super_gradients/common/registry/与src/super_gradients/common/factories/下的源码深入剖析其底层实现。读完本文你将掌握在 SuperGradients 中写名字即得对象的完整实战方案。阅读本文前建议先了解使用配置文件训练 与 训练配方Recipe入门。一、为什么需要 Factories配置与实例之间的鸿沟如果你翻阅过 SuperGradients 内置的 recipes 目录会发现大量对象直接在配方中定义。例如在 supervisely 数据集配方 中训练数据增强是这样声明的train_dataset_params: transforms: - SegColorJitter: brightness: 0.1 contrast: 0.1 saturation: 0.1 - SegRandomFlip: prob: 0.5 - SegRandomRescale: scales: [0.4, 1.6]如果把这个.yaml配方原样加载成 Python 字典得到的是{ train_dataset_params: { transforms: [ { SegColorJitter: { brightness: 0.1, contrast: 0.1, saturation: 0.1 } }, { SegRandomFlip: { prob: 0.5 } }, { SegRandomRescale: { scales: [0.4, 1.6] } } ] } }问题随之而来训练流程真正需要的是SegColorJitter、SegRandomFlip、SegRandomRescale这些类的实例对象而不是描述它们的配置字典。手写SegColorJitter(brightness0.1, contrast0.1, saturation0.1)这类实例化代码会随着配置项增多而迅速失控。Factories 正是为解决这一痛点而生上述对象在 SuperGradients 内部预先注册过因此当你在配方中写下这些名字时SuperGradients 会自动侦测到并替你完成实例化。二、使用内置工厂在 Recipe 中声明对象使用既有工厂不需要写任何额外代码。所有内置类在导入super_gradients包时即已完成注册配方中写下名字即可生效。以文档中的分割增强为例这三个类在源码中都有对应实现src/super_gradients/training/transforms/transforms.pySegRandomFliptransforms.py#L80-L100以概率prob同步水平翻转图像与掩码构造函数通过assert 0.0 prob 1.0校验概率取值范围默认值为0.5SegRandomRescaletransforms.py#L155-L198在scales给定的[min, max]区间内随机取缩放因子同步缩放图像与掩码。scales既支持(min, max)元组/列表也支持单个浮点数大于 1 时视为(1, scales)否则视为(scales, 1)check_valid_arguments还会自动纠正上下界顺序并对负数抛出ValueErrorSegColorJittertransforms.py#L394-L404对图像做颜色抖动亮度、对比度、饱和度、色相内部委托给torchvision.transforms.ColorJitter四个参数默认值均为0。这些类的注册名来自 src/super_gradients/common/object_names.py 中的枚举常量如Transforms.SegRandomFlip SegRandomFlip注册名与类名保持一致因此配方里写SegRandomFlip即可命中。仓库中真实的 supervisely 配方还展示了更完整的增强链supervisely_persons_dataset_params.yaml包含SegPadShortToCropSize、SegCropImageAndMask、SegStandardize、SegNormalize、SegConvertToTensor等验证集则使用确定性的SegResize——随机增强只用于训练集验证集用确定性变换这是值得沿用的实践。三、注册自定义类让配方认识你的对象如前所述只有被注册过的对象才能被实例化。注册的本质是把对象名字映射到对应的类类型这一过程记录下来。上文例子中字符串SegColorJitter被映射到了SegColorJitter类SuperGradients 因此知道如何把配方里的字符串变成对象。3.1 注册示例from super_gradients.common.registry import register_transform register_transform(nameMyTransformName) class MyTransform: def __init__(self, prob: float): ...这个例子把类MyTransform注册到了名字MyTransformName两者不同。虽然 SuperGradients 允许这样做但强烈建议不要这么做——注册名与类名不一致会降低可读性请始终用类自身的名字注册。3.2 在配方中使用自定义类注册完成后即可把新变换加入原有配方train_dataset_params: transforms: - SegColorJitter: brightness: 0.1 contrast: 0.1 saturation: 0.1 - SegRandomFlip: prob: 0.5 - SegRandomRescale: scales: [0.4, 1.6] - MyTransformName: # 这里使用注册名可以与类名不同 prob: 0.73.3 关键一步导入模块以触发注册注册发生在模块导入时。因此最后一步必须在你的训练脚本中导入包含MyTransformName的模块——导入动作会触发register_transform的执行SuperGradients 才能认识这个名字。以下脚本改编自 train_from_recipe 脚本from .my_module import MyTransform # 导入模块即可会触发 register_transform 函数 # 以下代码与基础的 train_from_recipe.py 完全一致 from omegaconf import DictConfig import hydra from super_gradients import Trainer, init_trainer hydra.main(config_pathrecipes, version_base1.2) def _main(cfg: DictConfig) - None: Trainer.train_from_config(cfg) def main() - None: init_trainer() # init_trainer 需要在 hydra.main 之前调用 _main() if __name__ __main__: main()仓库中 train_from_recipe_with_user_objects/train.py 就是这一模式的官方示例它在脚本顶部通过from ...user_dataset import user_mnist_train, user_mnist_val触发注册随后调用Trainer.train_from_config(cfg)。对应的 user_dataset.py 展示了完整闭环——用register_dataloader(user_mnist_train)注册自定义 dataloader用resolve_param让数据集接收配方中的变换配置配套配方 user_recipe_mnist_example.yaml 中只需写transforms: [RandomHorizontalFlip, ToTensor]即可这两个名字来自预注册的 torchvision 变换。3.4 注册机制的底层实现所有注册装饰器都由同一个工厂函数生成src/super_gradients/common/registry/registry.py#L14-L61def create_register_decorator(registry: Dict[str, Callable]) - Callable: def register(name: Optional[str] None, deprecated_name: Optional[str] None) - Callable: def decorator(cls: Callable) - Callable: def _registered_cls(registration_name: str): if registration_name in registry: registered_cls registry[registration_name] if registered_cls ! cls: raise Exception( f{registration_name} is already registered and points to {inspect.getmodule(registered_cls).__name__}.{registered_cls.__name__} ) registry[registration_name] cls registration_name name or cls.__name__ # 未指定 name 时使用类名注册 _registered_cls(registration_nameregistration_name) ... return cls return decorator return register从中可以读出三条重要规则注册名缺省时取cls.__name__这与建议用类名注册的官方建议一致同名重复注册且指向不同类会抛出Exception避免静默覆盖导致的隐患支持deprecated_name参数做兼容性注册旧名字照常可解析但会通过warn_if_deprecatedregistry.py#L64-L72发出DeprecationWarning提示用户改用新名字。四、底层原理Factory 如何把配置变成对象前面了解了如何用与如何注册这一节深入get()的调用链。4.1 基础用法直接调用工厂所有工厂的核心入口是get()最基础的用法如下from super_gradients.common.factories import TransformsFactory factory TransformsFactory() my_transform factory.get({MyTransformName: {prob: 0.7}})不难发现传给factory.get的入参正是配方加载后得到的那个字典参见本文第一节。工厂接收三类输入base_factory.py#L37-L72输入类型语义处理方式str类型名字查表后无参实例化self.type_dict[conf]()Mapping单键字典{类型名: {参数...}}查表后用参数字典实例化self.type_dict_type其他值已经是对象原样返回透传其中 Mapping 分支有两个值得注意的细节若字典包含多个键会抛出RuntimeError(Malformed object definition in configuration...)提示应为单条目的{类型名: {参数}}字典查找时先精确匹配未命中则走模糊匹配fuzzy_str会去除首尾空白、转小写并剔除下划线等符号见 utils.py#L246-L280让配置书写对大小写和下划线更宽容。若名字完全未知则抛出UnknownTypeExceptionfactory_exceptions.py——该异常会列出所有合法类型并借助 rapidfuzz 做模糊匹配给出提示Did you mean: ...?相似度得分大于 70 时对排查拼写错误非常友好。此外TransformsFactory在基类逻辑之外还做了特殊处理transforms_factory.py#L14-L28遇到Albumentations键时委托AlbumentationsTransformsFactory处理albumentations 库的类通过 albumentation.py 反射式批量注册并用AlbumentationsAdaptor包装为 SuperGradients 兼容的变换对 bbox/keypoint 参数有强约束format、label_fields等固定参数不可覆盖遇到Compose键时递归地用ListFactory解析其transforms子列表直接传入列表/ListConfig时交由ListFactory逐个元素实例化list_factory.py。4.2 推荐用法resolve_param装饰器工厂的真正威力在于与resolve_param装饰器配合使用factory_decorator.py#L11-L40。它让函数能够同时接受已实例化的对象或描述该对象的字典——也就是说你既可以传真实的 Python 对象也可以直接从配方传一个描述它的字典。class ImageNetDataset(torch_datasets.ImageFolder): resolve_param(transforms, factoryTransformsFactory()) def __init__(self, root: str, transform: Transform): ...此后ImageNetDataset既能接收MyTransform实例my_transform MyTransform(prob0.7) ImageNetDataset(root..., transformmy_transform)也能接收代表同一对象的字典my_transform {MyTransformName: {prob: 0.7}} ImageNetDataset(root..., transformmy_transform)第二种写法与.yaml配方概念完美契合——配方的dataset_params节点正是这种字典结构。resolve_param的实现会同时处理位置参数与关键字参数先在kwargs中查找参数名否则通过inspect.getfullargspec定位其在形参列表中的位置把对应值交给factory.get解析后再放回原位置。这一装饰器在仓库中被广泛使用例如 cifar.py、imagenet_dataset.py、detection_dataset.py、coco_keypoints.py、segmentation_dataset.py 等处均为数据集构造函数标注了resolve_param(transforms, ...)。register_transform与resolve_param的分工register_transform负责把字符串映射到类类型维护注册表resolve_param(transforms, factoryTransformsFactory())负责把配置转换成对象转换时借助前者的映射关系。一句话前者定义名字 → 类的字典后者消费这本字典完成实例化。五、SuperGradients 支持的工厂类型全景到目前为止我们只关注了TransformsFactory对应register_transform。实际上 SuperGradients 围绕训练全流程提供了丰富的工厂每个都有对应的注册装饰器全部导出自 src/super_gradients/common/factories/init.py 与 src/super_gradients/common/registry/init.pyfrom super_gradients.common.factories import ( register_model, register_kd_model, register_detection_module, register_metric, register_loss, register_dataloader, register_callback, register_transform, register_dataset, register_pre_launch_callback, register_unet_backbone_stage, register_unet_up_block, register_target_generator, register_lr_scheduler, register_lr_warmup, register_sg_logger, register_collate_function, register_sampler, register_optimizer, register_processing, )每个装饰器背后都对应一个注册表字典registry.py#L75-L196注册装饰器注册表典型用途register_modelARCHITECTURES网络架构如register_model(Models.DENSENET121)register_kd_modelKD_ARCHITECTURES知识蒸馏KD架构register_detection_moduleALL_DETECTION_MODULES检测专用模块NMS、anchor 生成器等register_metricMETRICS评估指标由MetricsFactory消费register_lossLOSSES损失函数内置MSE还注册了弃用名mse以兼容旧配置register_dataloaderALL_DATALOADERS数据加载器可为函数而非类register_callbackCALLBACKS训练回调由CallbacksFactory消费register_transformTRANSFORMS数据变换register_datasetALL_DATASETS数据集类register_pre_launch_callbackALL_PRE_LAUNCH_CALLBACKS启动前回调register_unet_backbone_stage/register_unet_up_blockBACKBONE_STAGES/UP_FUSE_BLOCKSUNet 骨干阶段与上采样块register_target_generatorALL_TARGET_GENERATORS检测目标生成器register_lr_scheduler/register_lr_warmupLR_SCHEDULERS_CLS_DICT/LR_WARMUP_CLS_DICT学习率调度与预热register_sg_loggerSG_LOGGERS实验日志记录器register_collate_functionALL_COLLATE_FUNCTIONS批处理 collate 函数register_samplerSAMPLERS数据采样器register_optimizerOPTIMIZERS优化器register_processingPROCESSINGS预处理模块部分注册表在模块加载时即被预填充了 PyTorch 生态的常用类无需二次注册即可在配方中使用TRANSFORMS预置了torchvision.transforms的Compose、ToTensor、Normalize、Resize、RandomCrop、ColorJitter等二十余个标准变换SAMPLERS预置了DistributedSampler、SequentialSampler、RandomSampler、WeightedRandomSampler等OPTIMIZERS预置了SGD、Adam、AdamW、RMSprop另有独立的TORCH_LR_SCHEDULERS字典收录StepLR、CosineAnnealingLR、ReduceLROnPlateau等 torch 调度器。此外type_factory.py 中的TypeFactory是工厂家族的特殊成员它返回类类型本身而非实例TypeFactory.from_enum_cls可由枚举一键构建并支持点分路径导入——配置中写package_name.sub_package.MyClass即可通过importlib动态导入任意类为使用未注册的外部类提供了兜底方案。六、总结与下一步围绕 Factories本指南覆盖了三个层次使用既有工厂SuperGradients 如何在配方中自动实例化对象——名字即配置配置即对象注册新类通过register_transform等装饰器将对象名映射到类类型在脚本顶部导入模块触发注册从而把自定义对象接入配方底层机制factory.get()对字符串/字典/列表三类输入的分派逻辑、resolve_param让函数同时接受对象与配置字典以及覆盖训练全流程的十余种工厂类型。Factories 是 SuperGradients 的核心组件之一它弥合了配置与实例化之间的缝隙让 YAML 配方既能被人类轻松阅读又能被机器直接执行。如果你打算构建自己的训练配方可以继续阅读下一篇教程 Recipes_Custom.md自定义 Recipe 实战学习如何从零组织一份可训练的配方文件。【免费下载链接】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),仅供参考