资讯动态

SuperGradients 损失函数(Loss)完全指南:内置损失、自定义损失与配置化训练

发布时间:2026/9/18 9:10:20 来源:尧图企业网站定制
SuperGradients 损失函数Loss完全指南内置损失、自定义损失与配置化训练【免费下载链接】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导读在 SuperGradients 中损失函数Loss是训练流水线的核心组件之一。无论是图像分类、目标检测YoloX / YOLO-NAS / SSD / PP-YOLOE、语义分割ShelfNet / STDC还是姿态估计SuperGradients 都为每种任务内置了经过验证的损失实现并支持直接使用任意 PyTorch 损失、通过criterion_params配置化传参、以及用register_loss注册自定义损失。读完本文你将掌握三种在训练中指定损失函数的方式字符串别名、配置字典、实例化对象理解loss与criterion_params的分工并能写出带细粒度日志与断点监控的自定义损失。本文内容以 Losses.md 为基础骨架并结合 losses 模块源码 与 registry 实现 进行展开。一、SuperGradients 损失机制概览1.1 一切损失都是 torch.nn.ModuleSuperGradients 对损失函数没有引入任何私有协议任何基于 PyTorch 的损失函数都可以直接使用。内置的多任务损失实现同样只是torch.nn.Module的子类文档与配置中出现的CrossEntropyLoss、YoloXDetectionLoss等字符串本质上都是这些类在注册表中的别名string alias。这一点在 losses/init.py 中得到直接印证所有内置损失均以类形式导出例如from super_gradients.training.losses.yolox_loss import YoloXDetectionLoss, YoloXFastDetectionLoss from super_gradients.training.losses.ssd_loss import SSDLoss from super_gradients.training.losses.bce_dice_loss import BCEDiceLoss1.2 内置损失一览原文档列出的全部内置损失如下均在 registry.py 中通过register_loss注册损失名称字符串别名对应类 / 说明CrossEntropyLoss带标签平滑的交叉熵分类损失见 label_smoothing_cross_entropy_loss.pyMSE均方误差注册时直接映射到torch.nn.MSELoss见 registry.pyRSquaredLossR² 回归损失见 r_squared_loss.pyShelfNetOHEMLossShelfNet 语义分割在线难例挖掘损失见 shelfnet_ohem_loss.pyShelfNetSemanticEncodingLossShelfNet 语义编码损失见 shelfnet_semantic_encoding_loss.pyYoloXDetectionLossYoloX 检测损失objectness IoU 分类 L1 四个分量见 yolox_loss.pyYoloXFastDetectionLossYoloX 检测损失的快速版本SSDLossSSD 检测损失见 ssd_loss.pySTDCLossSTDC 分割损失见 stdc_loss.pyBCEDiceLossBCE Dice 组合分割损失见 bce_dice_loss.pyKDLogitsLoss知识蒸馏 Logits 损失见 kd_losses.pyDiceCEEdgeLoss带边缘感知的 Dice CE 分割损失见 dice_ce_edge_loss.py此外从 losses/init.py 可以看到仓库还注册了PPYoloELossPP-YOLOE 检测、DEKRLossDEKR 姿态估计、RescoringLoss姿态重打分与YoloNASPoseLossYOLO-NAS-Pose 姿态等面向更新任务的损失说明注册表是持续扩展的。1.3 名称不区分大小写、不区分符号SuperGradients 的对象名解析是大小写不敏感、符号不敏感的传CrossEntropy、crossentropyloss等变体都可以命中CrossEntropyLoss。这一行为由 base_factory.py 中的fuzzy_str/fuzzy_keys模糊匹配逻辑保证先做精确匹配失败后做模糊匹配。二、内置损失的基本用法2.1 直接调用 Trainer.train(...)Python 脚本方式在my_training_script.py中通过training_params的loss键指定损失名称from super_gradients import Trainer trainer Trainer(external_criterion_test) train_dataloader ... valid_dataloader ... model ... train_params { max_epochs: 100, loss: CrossEntropyLoss, criterion_params: {}, # ... 其他训练参数优化器、学习率、指标等 } trainer.train(modelmodel, training_paramstrain_params, train_loadertrain_dataloader, valid_loadervalid_dataloader)由于名称不区分大小写与符号CrossEntropy同样合法。2.2 使用 object_names 模块享受 IDE 补全为避免手写字符串出错推荐使用 object_names.py 中定义的静态常量类Lossesfrom super_gradients.common.object_names import Losses train_params { ... loss: Losses.CROSS_ENTROPY, # 等价于 CrossEntropyLoss criterion_params: {}, ... }该常量类覆盖了全部注册损失包括Losses.MSE、Losses.YOLOX_LOSS、Losses.SSD_LOSS、Losses.KD_LOSS、Losses.BCE_DICE_LOSS、Losses.YOLONAS_POSE_LOSS等几乎所有主流 IDE 都能对其进行自动补全。2.3 配置文件方式train_from_recipe当使用train_from_recipe或任何最终调用Trainer.train_from_config(...)的入口如 train_from_recipe.py时在my_training_hyperparams.yaml中配置loss: YoloXDetectionLoss criterion_params: strides: [8, 16, 32] # 所有 yolo 输出的下采样步长 num_classes: 80 # 数据集类别数这里两个training_params参数各司其职loss定义损失的类型注册表中的字符串别名criterion_params一个字典会被解包unpack成底层YoloXDetectionLoss类的构造参数**criterion_params。从 yolox_loss.py 的构造签名可以看到YoloXDetectionLoss实际还支持更多可选参数参数默认值含义strides必填各 yolo 输出层的下采样步长如[8, 16, 32]num_classes必填检测类别数use_l1False是否在非增强阶段使用 L1 回归损失center_sampling_radius2.5正样本中心采样半径iou_typeiouIoU 计算方式iou_weight5.0IoU 损失权重obj_weight1.0objectness 损失权重cls_weight1.0分类损失权重cls_pos_weightNone分类正样本权重dynamic_ks_bias1.1动态 top-k 正样本分配偏置sync_num_fgsFalse是否在 DDP 训练下同步正样本数量obj_loss_fixFalse是否按总 anchor 数而非匹配正样本数计算 objectness 损失该损失由四个分量构成L L_objectness L_iou L_classification 1[no_aug_epoch] * L_l1详见 yolox_loss.py 的注释。所有未在criterion_params中出现的参数均使用默认值这就是配置驱动训练的关键所在。2.4 底层解析LossesFactory 与 BaseFactorylosscriterion_params的解析链路为sg_trainer.py 中self.criterion LossesFactory().get({self.training_params.loss: self.training_params.criterion_params})。LossesFactory直接继承BaseFactory并以注册表LOSSES为类型字典见 losses_factory.py而BaseFactory.get见 base_factory.py的逻辑是若传入字符串在LOSSES注册表中精确/模糊查找无参实例化若传入{type_name: {params}}单元素字典以**params方式调用构造函数若传入其他对象如已实例化的nn.Module原样返回不经过注册表。这也解释了为什么criterion_params与loss必须配对使用以及传入实例化对象时criterion_params会被忽略见下文。三、直接传入实例化的 nn.Module 作为损失3.1 Trainer.train 脚本方式SuperGradients 同样支持直接传入已经实例化的nn.Module对象。在my_training_script.py中import torch trainer Trainer(external_criterion_test) train_dataloader ... valid_dataloader ... model ... train_params { ... loss: torch.nn.CrossEntropyLoss(), # 直接传入实例对象 ... } trainer.train(modelmodel, training_paramstrain_params, train_loadertrain_dataloader, valid_loadervalid_dataloader)3.2 配置文件方式Hydra 实例化使用配置驱动时可以用_target_语法直接实例化任意类该机制由 Hydra 的实例化功能提供。在my_training_hyperparams.yaml中loss: _target_: torch.nn.CrossEntropyLoss注意当loss传入的是实例化对象时criterion_params将被忽略——因为构造参数已经由_target_节点的配置决定了。这一方式虽然没有使用register_loss那样便捷但在需要快速试用一个未注册的 PyTorch 原生损失时非常实用。四、使用自定义损失Using Your Own Loss4.1 基本要求与 forward 签名SuperGradients 支持用户自定义损失前提是继承torch.nn.Moduleforward签名为forward(preds, target)形式——第一个参数是模型输出第二个参数是标签/真值参数名任意不要求必须叫preds/target。import torch.nn as nn class MyLoss(nn.Module): ... def forward(self, preds, target): ...目前forward仅接受两个位置参数未来版本会支持接受额外参数的自定义损失。4.2 日志输出机制loss 与 loss_items在最常见场景下损失函数对每个 batch 返回一个标量张量用于反向传播训练过程中它会以LOSS_CLASS.__name__为键被记录到日志TensorBoard 及所有受支持的 SGLogger 中并按 epoch 聚合展示。为了记录更细粒度的信息forward(...)也可以返回(loss, loss_items)元组loss用于反向传播的张量即原始损失输出loss_items形状为(n_items)的张量包含前向过程中计算出的、希望在整个 epoch 上记录的各个分量。典型的场景是总损失由若干分量求和得到而你希望把每个分量都记进日志。示例class MyLoss(nn.Module): ... def forward(self, inputs, targets): ... total_loss comp1 comp2 loss_items torch.cat((total_loss.unsqueeze(0), comp1.unsqueeze(0), comp2.unsqueeze(0))).detach() return total_loss, loss_items train_params { ..., loss: MyLoss(), metric_to_watch: MyLoss/loss_0 } Trainer.train(..., train_paramstrain_params)上述代码会按loss_items的位置下标记录MyLoss/loss_0、MyLoss/loss_1、MyLoss/loss_2三个值同时由于metric_to_watch设为MyLoss/loss_0每当该指标达到历史最优时就会保存模型检查点ckpt_best.pth。4.3 用 component_names 给分量命名为了可读性可以在损失类中定义component_names属性——一个长度为n_items的字符串列表其第 i 个元素对应loss_items第 i 个分量的名称。此后每个分量都会以LOSS_CLASS.__name__/COMPONENT_NAME的形式被记录、渲染到 TensorBoard并可作为metric_to_watch的监控对象class MyLoss(nn.Module): ... def forward(self, inputs, targets): ... total_loss comp1 comp2 loss_items torch.cat((total_loss.unsqueeze(0), comp1.unsqueeze(0), comp2.unsqueeze(0))).detach() return total_loss, loss_items property def component_names(self): return [total_loss, my_1st_component, my_2nd_component] train_params { ..., loss: MyLoss(), metric_to_watch: MyLoss/my_1st_component } Trainer.train(..., train_paramstrain_params)此时日志与监控项为MyLoss/total_loss、MyLoss/my_1st_component、MyLoss/my_2nd_component。需要强调的是由于运行日志会把loss_items保存在内部状态中用于 epoch 级聚合强烈建议对loss_items调用.detach()脱离计算图以节省显存/内存。这一点在原文档与 sg_trainer.py 的说明中均有明确提示。此外loss_logging_items_names这一training_params参数也提供了另一种为损失输出命名的途径见 sg_trainer.py。4.4 配置文件中使用自定义损失register_loss当使用配置文件训练train_from_recipe或Trainer.train_from_config(...)时需要通过register_loss装饰器把自定义损失注册进注册表。第一步在my_loss.py中定义并注册损失类import torch.nn as nn from super_gradients.common.registry import register_loss register_loss(my_loss) class MyLoss(nn.Module): ...register_loss是 registry.py 中通过create_register_decorator(registryLOSSES)生成的装饰器工厂。它支持两种注册名name参数指定正式注册名deprecated_name参数则额外注册一个已废弃别名用于兼容旧配置并在使用该旧名时给出DeprecationWarning见 registry.py。若省略name则默认以类名注册。第二步在my_training_hyperparams.yaml中像使用内置损失一样使用它loss: my_loss criterion_params: ... # 这些参数会解包给 MyLoss 的构造函数第三步在入口脚本中导入该模块以触发注册。虽然类本身没有被直接使用但导入动作会执行装饰器、把类写入LOSSES注册表from omegaconf import DictConfig import hydra import pkg_resources from my_loss import MyLoss # 仅用于触发 register_loss 注册 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()注册完成后my_loss就与CrossEntropyLoss等内置损失完全等价既支持字符串别名解析也支持criterion_params构造参数解包整个解析链路复用同一套LossesFactory/BaseFactory机制。五、实战场景损失在真实 Recipe 中的形态仓库的 recipe 配置是理解losscriterion_params用法的最佳示例recipes 目录分类场景如 imagenet_resnet50.yaml使用CrossEntropyLoss检测场景如 coco2017_yolo_nas_s.yaml、coco2017_ppyoloe_s.yaml使用YoloXDetectionLoss/PPYoloELoss并配套strides、num_classes等criterion_params分割场景如 cityscapes_ddrnet.yaml使用CrossEntropyLoss或其变体知识蒸馏场景如 imagenet_resnet50_kd.yaml使用KDLogitsLoss其构造参数包含task_loss_fn任务损失、distillation_loss_fn蒸馏损失默认KDklDivLoss与distillation_loss_coeff蒸馏损失系数默认 0.5见 kd_losses.py。在切换到自定义数据与任务时只需修改 recipe 中的loss与criterion_params即可在不改动训练代码的情况下完成损失更换这正是 SuperGradients 配置驱动训练的核心价值。六、总结与要点速查使用方式适用入口loss 的写法criterion_params 是否生效字符串别名Trainer.train/ recipeCrossEntropyLoss、Losses.CROSS_ENTROPY生效解包为构造参数配置字典recipeHydra{_target_: torch.nn.CrossEntropyLoss}忽略实例化对象Trainer.traintorch.nn.CrossEntropyLoss()忽略无实例化环节自定义损失两者皆可register_loss(my_loss)后使用my_loss生效关键结论回顾SuperGradients 的损失体系完全建立在torch.nn.Module之上内置损失名只是注册表别名loss定义类型、criterion_params提供构造参数二者通过LossesFactorylosses_factory.py统一解析名称匹配不区分大小写与符号自定义损失需遵循forward(preds, target)签名返回(loss, loss_items)元组可让各分量按component_names或位置下标被日志、TensorBoard 与metric_to_watch断点监控记得对loss_items调用.detach()配置驱动场景下使用register_loss注册的损失与内置损失完全等价仅需在入口脚本中导入一次以触发注册。至此你已经掌握了在 SuperGradients 中配置、扩展与调试损失函数的全部主流路径可以据此在自己的分类、检测、分割或蒸馏任务中自由选用与定制损失。【免费下载链接】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),仅供参考

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

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

免费获取报价