
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导读模型检查点Checkpoint是深度学习训练流程中可以回退的锚点它既记录了模型在训练各阶段的状态快照也承载了中断后无缝恢复训练的能力。本文以 SuperGradientsSG官方文档documentation/source/Checkpoints.md为主体结合仓库内 Trainer 与 checkpoint 工具源码系统讲解 SG 中检查点的保存时机与目录组织、内部数据结构、严格/宽松/按形状匹配等加载策略、基于 recipe 的恢复训练与远程WandB断点续训以及一键评估历史检查点的方法。读完本文你将能够熟练运用models.get、load_checkpoint_to_model、Trainer.resume_experiment等 API独立完成从保存到加载再到恢复、评估的完整闭环。什么是模型检查点训练过程中模型性能会随着它看到的样本数量不断变化。业界最佳实践是在训练的关键节点保存模型状态——每一份保存下来的状态就是一个 checkpoint对应模型开发过程中某个时刻的版本。训练结束后应当使用验证集上表现最佳的 checkpoint 作为最终产物同时checkpoint 也保证了训练被中断时可以从断点继续而不必从头再来。SuperGradients 遵循这一思想在训练全程自动保存多个不同用途的 checkpoint。每个 checkpoint 都对应一个训练阶段彼此分工明确有的用于追踪最佳性能有的用于支持断点续训有的用于产出最终部署模型。如果你还不熟悉 SG 的实验管理机制建议先阅读 实验管理文档 了解ckpt_root_dir、experiment_name、run等概念。Checkpoint 保存策略哪个文件、何时保存、存在哪里四种默认检查点文件在 SG 中训练过程中会按照下表所示的时机自动保存不同类型的 checkpoint 文件Checkpoint 文件名保存时机ckpt_best.pth每次验证时达到新的最佳metric_to_watch指标即覆盖保存ckpt_latest.pth每个 epoch 结束时保存持续覆盖为最新状态average_model.pth训练结束时保存由验证指标最优的 10 个模型快照平均得到仅当average_best_modelsTrue时生成ckpt_epoch_{EPOCH_INDEX}.pth当save_ckpt_epoch_list训练参数中指定了固定 epoch 序号EPOCH_INDEX时在该 epoch 结束时保存其中几个关键训练参数在 默认训练超参配置 中有明确定义metric_to_watch决定最佳模型评判依据的验证指标默认Accuracy是ckpt_best.pth是否被覆盖的标尺greater_metric_to_watch_is_better为True时最大化该指标为最佳为False时最小化如损失类指标save_ckpt_epoch_list需要额外保存的 epoch 序号列表例如[10, 15]average_best_models是否在训练结束时保存平均模型ckpt_best_name最佳模型的输出文件名默认ckpt_best.pthsave_model总开关控制是否保存模型检查点。从源码看这些逻辑集中在 sg_trainer.py 的_save_checkpoint中每次验证结束后先根据metric_to_watch与greater_metric_to_watch_is_better判断当前模型是否刷新了best_metric随后统一构建 state dict先写ckpt_latest.pth再按save_ckpt_epoch_list写定点 epoch 文件若指标更优则覆盖ckpt_best.pth最后若开启average_best_models则计算平均模型。而10 个最优快照的滑动平均由 weight_averaging_utils.py 的get_average_model实现它维护一个按验证指标排序的快照池每当新模型优于池中最差者就替换之训练结束时对池内快照做逐层加权平均。检查点保存位置与目录结构检查点文件统一保存在如下路径中ckpt_root_dir/experiment_name/run_dir其中ckpt_root_dir与experiment_name由用户在实例化Trainer时指定Trainer(ckpt_root_dirpath/to/ckpt_root_dir, experiment_namemy_experiment)run_dir是每次调用trainer.train(...)启动新一轮训练时自动生成的唯一目录形如RUN_20230802_131052_651906保证同一实验下的多次运行互不覆盖。当使用克隆下来的 SuperGradients 仓库源码直接运行时可以省略ckpt_root_dir参数此时检查点默认保存到仓库的super_gradients/checkpoints目录下。一次典型训练结束后ckpt_root_dir下的完整结构如下ckpt_root_dir │ ├── experiment_name │ │ │ ├─── run_dir │ │ ├─ ckpt_best.pth # 验证集上表现最佳的检查点 │ │ ├─ ckpt_latest.pth # 最近一个 epoch 结束时的检查点 │ │ ├─ average_model.pth # 最优模型快照的平均模型 │ │ ├─ ckpt_epoch_*.pth # 指定 epoch 的检查点如 epoch 10、15 │ │ ├─ events.out.tfevents.* # TensorFlow 运行产物 │ │ └─ log_timestamp.txt # 本次运行的 Trainer 日志 │ │ │ └─── other_run_dir │ └─ ... │ └─── other_experiment_name │ ├─── run_dir │ └─ ... │ └─── another_run_dir └─ ...这一实验 → 运行 → 文件三层组织方式配合 SG 的日志体系使得同一实验名下的多个 run 天然隔离也便于后续通过run_id精准定位任意一次运行的检查点。Checkpoint 内部结构SG 的检查点是 PyTorchstate_dict的实例参见 PyTorch 官方关于 state_dict 的说明除了模型权重之外还携带了训练相关的附加信息因此可以支撑仅加载权重与完整恢复训练两类场景。检查点的顶层键key如下键名含义net网络的state_dict模型权重acc网络在验证集上取得的metric_to_watch指标值floatepoch最近完成的一个 epochoptimizer_state_dict优化器的state_dictscaler_state_dict可选——仅当以mixed_precisionTrue训练时存在为Trainer.scaler的state_dictema_net可选——仅当以emaTrue训练时存在为 EMA 模型的state_dict。注意average_model.pth即使开启 EMA 也不含该键因为平均模型本身就是对 EMA 快照做平均其net键已经是 EMA 快照的平均结果torch_scheduler_state_dict可选——仅当使用 PyTorch 原生 LR scheduler 时存在见 LRScheduling对照源码_save_checkpoint构建 state dict 时除上述键外还额外写入metrics本次验证及训练所有指标的汇总 dictpackages训练环境已安装的包及版本列表便于复现环境_best_ckpt_metrics最佳检查点对应的指标记录processing_params由验证数据加载器推导出的数据预处理参数便于加载后直接做推理预测。其中ema_net通过unwrap_model(self.ema_model.ema).state_dict()提取scaler_state_dict仅在self.scaler is not None时写入torch_scheduler_state_dict仅在使用了 torch 原生 scheduler 时写入且内部通过get_scheduler_state做了与 PyTorch 版本的兼容处理见 checkpoint_utils.py 的get_scheduler_state。相关概念可进一步参考 混合精度训练文档 与 EMA 文档。通过 SG Logger 远程保存检查点SG 支持借助第三方实验追踪工具远程保存检查点例如 Weights Biases。只需在sg_logger_params训练参数中设置save_checkpoints_remoteTrue训练过程中产出的 checkpoint 就会被同步上传到远程存储。更完整的配置说明见 第三方实验监控文档。远程保存的意义不仅在于备份它还是训练中断后从云端续训的前提具体操作见下文从远程存储恢复训练一节。加载检查点加载权重与恢复训练是两回事加载检查点可以按使用场景拆分为两类仅加载模型权重用于推理、迁移学习、微调完整恢复训练连同优化器、scheduler、scaler 状态一起恢复。后者需要的状态信息更多SG 在Trainer内部负责将其组装还原前者则只需要net权重。SG 的加载方法相比 PyTorch 原生load_state_dict()提供了更多能力尤其是针对 SG 自身训练的检查点。通过 models.get 加载权重权重加载可以直接在模型初始化后完成有两种等价途径在models.get(...)中传入checkpoint_path或对已有的torch.nn.Module实例显式调用load_checkpoint_to_model。假设我们启动过一次与下面结构类似的训练实验from super_gradients.training import Trainer ... ... from super_gradients.training import models from super_gradients.common.object_names import Models trainer Trainer(my_resnet18_training_experiment, ckpt_root_dir/path/to/my_checkpoints_folder) train_dataloader ... valid_dataloader ... model models.get(model_nameModels.RESNET18, num_classes10) train_params { ... loss: CrossEntropyLoss, criterion_params: {}, save_ckpt_epoch_list: [10, 15] ... } trainer.train(modelmodel, training_paramstrain_params, train_loadertrain_dataloader, valid_loadervalid_dataloader)训练结束后我们想加载ckpt_best.pth的权重只需把它的完整路径传给models.get的checkpoint_path参数from super_gradients.training import models from super_gradients.common.object_names import Models model models.get( model_nameModels.RESNET18, num_classes10, checkpoint_path/path/to/my_checkpoints_folder/my_resnet18_training_experiment/RUN_20230802_131052_651906/ckpt_best.pth, )重要提示通过models.get(...)加载 SG 训练的检查点时如果网络是以 EMA 方式训练的默认加载的是 EMA 权重。这与源码中_load_weights的行为一致若 checkpoint 中存在ema_net会先将其替换为net再加载见 checkpoint_utils.py。如果已经持有模型实例也可以直接使用load_checkpoint_to_modelfrom super_gradients.training import models from super_gradients.common.object_names import Models from super_gradients.training.utils.checkpoint_utils import load_checkpoint_to_model model models.get(model_nameModels.RESNET18, num_classes10) load_checkpoint_to_model( netmodel, ckpt_local_path/path/to/my_checkpoints_folder/my_resnet18_training_experiment/RUN_20230802_131052_651906/ckpt_best.pth, )从源码看load_checkpoint_to_model的完整签名还支持更多能力load_backbone仅将权重加载到模型的backbone子模块要求模型具备backbone属性strict加载的键匹配严格度默认NO_KEY_MATCHING见下一节load_weights_only加载后丢弃net以外的所有附加信息load_ema_as_net显式要求加载ema_net作为网络权重不存在时会抛错load_processing_params是否将 checkpoint 内的processing_params应用到模型的set_dataset_processing_params。底层实现上load_checkpoint_to_model会先通过read_ckpt_state_dict读取 checkpoint支持本地路径与https://URL见 checkpoint_utils.py再交给adaptive_load_state_dict完成实际装载。StrictLoad扩展 PyTorch 的 strict 参数如果不熟悉 PyTorchload_state_dict()的strict参数语义建议先阅读 PyTorch 官方保存与加载模型教程。SG 中models.get()和load_checkpoint_to_model分别用strict与strict_load两个参数承担 PyTorchstrict的职责但它们接受的是 SG 自定义的StrictLoad枚举类型。该枚举定义于 strict_load.pyclass StrictLoad(Enum): Wrapper for adding more functionality to torchs strict_load parameter in load_state_dict(). Attributes: OFF - Native torch strict_load off behavior. See nn.Module.load_state_dict() documentation for more details. ON - Native torch strict_load on behavior. See nn.Module.load_state_dict() documentation for more details. NO_KEY_MATCHING - Allows the usage of SuperGradients adapt_checkpoint function, which loads a checkpoint by matching each layers shapes (and bypasses the strict matching of the names of each layer (i.e., disregards the state_dict key matching)). KEY_MATCHING - Loose load strategy that loads the state dict from checkpoint into model only for common keys and also handles the case when shapes of the tensors in the state dict and model are different for the same key (Such layers will be skipped). OFF False ON True NO_KEY_MATCHING no_key_matching KEY_MATCHING key_matching也就是说除了 PyTorch 原生语义的OFF等价strictFalse与ON等价strictTrueSG 还额外提供了两种宽松加载模式NO_KEY_MATCHING利用state_dict是OrderedDict这一事实按层顺序做形状匹配加载完全忽略键名匹配。当网络底层结构一致、但各层的state_dict键名与模型内键名不一致时非常有用。从源码看adaptive_load_state_dict会先尝试按 strict 加载失败后再调用adapt_state_dict_to_fit_model_layer_names将 checkpoint 键名重排为模型键名最终以strictTrue完成加载。KEY_MATCHING仅加载 checkpoint 与模型共有的键且同一键下张量形状不一致的层会被跳过。其实现是 transfer_weights逐个键尝试load_state_dict(..., strictFalse)形状不兼容的层直接跳过。下面用一个简单例子演示不同 strict 模式的行为差异import torch class ModelA(torch.nn.Module): def __init__(self): super(ModelA, self).__init__() self.conv1 torch.nn.Conv2d(3, 6, 5) self.conv2 torch.nn.Conv2d(6, 16, 5) class ModelB(torch.nn.Module): def __init__(self): super(ModelB, self).__init__() self.conv1 torch.nn.Conv2d(3, 6, 5) self.CONV2 torch.nn.Sequential([torch.nn.Conv2d(6, 16, 5)])上述两个网络的权重结构完全一致但state_dict的键名不同conv2vsCONV2.0。因此使用strictTrue即StrictLoad.ON从一个加载到另一个会直接报错崩溃使用strictFalse即StrictLoad.OFF不会崩溃但只能成功加载第一个卷积层的权重使用 SG 的no_key_matching即StrictLoad.NO_KEY_MATCHING可以完整、正确地完成两者之间的权重迁移。另外从源码还可以看到adaptive_load_state_dict对旧版本 checkpoint 有向后兼容处理若所有键都以module.前缀开头即由 DataParallel/DistributedDataParallel 包装保存的 checkpoint会先通过 maybe_remove_module_prefix 自动去除该前缀再执行加载。加载 Model Zoo 的预训练权重通过models.get(...)三行代码即可加载任意 SG 预训练模型from super_gradients.training import models from super_gradients.common.object_names import Models model models.get(Models.YOLOX_S, pretrained_weightscoco)pretrained_weights参数指明预训练权重所用的数据集例如coco、imagenet。从源码看load_pretrained_weights会按architecture _ pretrained_weights在 pretrained_models.py 的MODEL_URLS字典中查找下载地址未命中则抛出MissingPretrainedWeightsException下载后同样走adaptive_load_state_dict并以StrictLoad.NO_KEY_MATCHING加载。值得注意的是预训练权重同样遵循优先加载 EMA 权重的规则对 YOLOX 系列源码使用专门的 YoloXCheckpointSolver 处理新旧版本键名差异其中layers_rename_table是代码生成的映射表并配有tests/unit_tests/yolox_unit_test.py中的单元测试验证对 YOLO-NAS 及 YOLO-NAS-POSE 系列下载时会输出许可证提示相关条款见仓库根目录的 LICENSE.YOLONAS.md 与 LICENSE.YOLONAS-POSE.md。通过配置文件加载检查点在基于配置文件的训练流程中检查点加载参数集中在checkpoint_params配置段。仓库中的默认模板位于 checkpoint_params/default_checkpoint_params.yaml其结构如下load_checkpoint: False # whether to load checkpoint load_backbone: False # whether to load only backbone part of checkpoint checkpoint_path: # checkpoint path that is located in super_gradients/checkpoints external_checkpoint_path: # checkpoint path that is not located in super_gradients/checkpoints source_ckpt_folder_name: # dirname for checkpoint loading strict_load: # key matching strictness for loading checkpoints weights _target_: super_gradients.training.sg_trainer.StrictLoad value: no_key_matching pretrained_weights: # a string describing the dataset of the pretrained weights (for example imagenent). # num_classes of checkpoint_path/ pretrained_weights, when checkpoint_path is not None. # Used when num_classes ! checkpoint_num_class. # In this case, the module will be initialized with checkpoint_num_class, then weights will be loaded. # Finally model.replace_head(new_num_classesnum_classes) is called to replace the head with new_num_classes. checkpoint_num_classes: # number of classes in the checkpoint这些参数正是用于以不同权重启动训练如微调场景——在Trainer.train_from_config(...)的底层流程中它们会被透传给models.get(...)classmethod def train_from_config(cls, cfg: Union[DictConfig, dict]) - Tuple[nn.Module, Tuple]: ... # BUILD NETWORK model models.get( ... strict_loadcfg.checkpoint_params.strict_load, pretrained_weightscfg.checkpoint_params.pretrained_weights, checkpoint_pathcfg.checkpoint_params.checkpoint_path, load_backbonecfg.checkpoint_params.load_backbone, ) # INSTANTIATE DATA LOADERS train_dataloader ... val_dataloader ... ... # TRAIN res trainer.train(...) ...对其中几个参数做进一步说明strict_load默认值为no_key_matching即默认以形状匹配的宽松方式加载最大化兼容 SG 训练产出的检查点checkpoint_num_classes用于换头微调当加载的 checkpoint 分类数与目标模型不同时SG 会先用 checkpoint 的类别数初始化模型并加载权重再调用model.replace_head(new_num_classesnum_classes)替换输出头避免因 head 维度不匹配而加载失败external_checkpoint_path与source_ckpt_folder_name用于加载不在默认 checkpoints 目录下的外部权重。恢复训练Resume TrainingSG 的断点续训由三个训练参数协同控制它们都在 default_train_params.yaml 中有定义提供了从最新断点到任意指定检查点分支的灵活度resume: False # 是否从同一实验名下的最新一次运行继续训练 run_id: # 同一实验内要从哪一次运行run恢复 resume_path: # 直接指定一个 .pth 检查点文件的路径来恢复训练1. 恢复最近一次运行将resumeTrue设置为真SG 会在同一实验名下找到最近一次运行的ckpt_latest.pth默认文件名由ckpt_name控制并从该断点继续# 从 cifar_experiment 最近一次运行处继续 python -m super_gradients.train_from_recipe --config-namecifar10_resnet experiment_namecifar_experiment training_hyperparams.resumeTrue2. 恢复指定运行通过run_id可以精确恢复同一实验内的某次运行# 从 cifar_experiment 中由 run_id 标识的那次运行继续 python -m super_gradients.train_from_recipe --config-namecifar10_resnet experiment_namecifar_experiment run_idRUN_20230802_131052_6519063. 从指定检查点分支通过resume_path指定任意.pth文件SG 会新建一个 run 目录从该检查点继续训练并把新产生的检查点保存到新目录中——这非常适合做分支实验# 从指定检查点分支创建一次新的 run python -m super_gradients.train_from_recipe --config-namecifar10_resnet experiment_namecifar_experiment training_hyperparams.resume_path/path/to/checkpoint.pth从源码看_load_checkpoint_to_model会综合resume、run_id、resume_path、resume_from_remote_sg_logger四个来源决定是否加载检查点其中resumeTrue且没有显式路径时沿用原 run而使用resume_from_remote_sg_logger或resume_path时则会生成新的 run_id。恢复时是否连同优化器状态一起加载由load_opt_params参数控制。4. 使用原始 recipe 恢复训练恢复训练是参数强相关的如果当前 recipe 定义的模型架构与 checkpoint 中的架构不一致就无法恢复训练加载时会直接抛出异常。典型场景是模型训练于一段时间之前期间你修改了模型架构定义此时再拿旧 checkpoint 恢复就会失败。为避免这一问题SG 提供了基于原始训练 recipe的恢复方式Trainer.resume_experiment(ckpt_root_dir..., experiment_name..., run_id...)run_id可选指定要恢复的具体运行缺省时自动恢复该实验最近一次运行内部通过get_latest_run_id定位。注意Trainer.resume_experiment只能恢复通过Trainer.train_from_config启动的训练因为恢复依赖训练时保存下来的完整配置快照sg_trainer.py 的 resume_experiment 会先读取历史 config再注入resumeTrue与run_id后重新走train_from_config。命令行等价用法可参考 resume_experiment.py 入口旧示例脚本位于 examples/resume_experiment_example/resume_experiment.py已标记弃用python -m super_gradients.resume_experiment --experiment_namemy_experiment_name从远程存储恢复训练WandBSG 支持从 SG Logger 定义的远程存储中恢复训练。前提是训练期间已在sg_logger_params中开启save_checkpoints_remoteTrue使检查点被同步到远程例如 WandB run 的存储。假设我们使用 WandB SG Logger 运行实验则training_hyperparams应包含sg_logger: wandb_sg_logger, # WeightsBiases Logger, see class super_gradients.common.sg_loggers.wandb_sg_logger.WandBSGLogger for details sg_logger_params: # Params that will be passes to __init__ of the logger super_gradients.common.sg_loggers.wandb_sg_logger.WandBSGLogger project_name: project_name, # WB project name save_checkpoints_remote: True, save_tensorboard_remote: True, save_logs_remote: True, entity: YOUR-ENTITY-NAME, # username or team name where youre sending runs api_server: OPTIONAL-WANDB-URL # Optional: In case your experiment tracking is not hosted at wandb serverssave_checkpoints_remoteTrue会促使训练全程在 WandB 中保存检查点。若此时训练被中断只需设置两个训练超参即可从 WandB run 存储中的检查点恢复设置resume_from_remote_sg_loggerresume_from_remote_sg_logger: True在sg_logger_params中通过wandb_id传入原 run 的 idsg_logger: wandb_sg_logger, # WeightsBiases Logger, see class super_gradients.common.sg_loggers.wandb_sg_logger.WandBSGLogger for details sg_logger_params: # Params that will be passes to __init__ of the logger super_gradients.common.sg_loggers.wandb_sg_logger.WandBSGLogger wandb_id: YOUR_RUN_ID project_name: project_name, # WB project name save_checkpoints_remote: True, save_tensorboard_remote: True, save_logs_remote: True, entity: YOUR-ENTITY-NAME, # username or team name where youre sending runs api_server: OPTIONAL-WANDB-URL # Optional: In case your experiment tracking is not hosted at wandb servers完成以上两步后重新启动训练ckpt_latest.pth默认可通过ckpt_name修改会被自动下载到本地 checkpoints 目录随后从该检查点继续训练——与本地断点续训完全一致。底层实现可见 default_train_params.yaml 中resume_from_remote_sg_logger的注释该机制目前仅支持 WandB Logger且仅对以save_checkpoints_remoteTrue运行的实验有效。评估检查点与用原始 recipe 恢复训练的思路类似我们常常希望在不重新熟悉训练配置的情况下直接评估某个历史检查点。为此 SG 提供了两个 Trainer 方法Trainer.evaluate_checkpoint(...)评估你自己此前某个实验产生的检查点使用该实验训练时完全一致的参数数据集、验证指标等。即便之后 recipe 被修改过评估仍按训练时的参数进行确保验证结果与训练时完全可比。Trainer.evaluate_recipe(...)当前源码中已演进为evaluate_from_config评估 Model Zoo 的预训练模型检查点或以不同的参数评估检查点例如更换数据集或验证指标。从源码看evaluate_checkpoint的实现逻辑是定位到指定实验默认最近 run的历史配置快照注入resumeTrue与目标ckpt_name然后调用evaluate_from_config而evaluate_from_config会实例化配置中的模型与验证数据加载器通过models.get(...)加载指定检查点后执行一次验证。注意旧方法名evaluate_from_recipe已从 3.6.2 起标记弃用并指向evaluate_from_config见 sg_trainer.py。命令行用法示例见 evaluate_checkpoint.py 入口旧示例脚本位于 examples/evaluate_checkpoint_example/evaluate_checkpoint.py同样已标记弃用# 评估实验 my_experiment_name 中的 average_model.pth python -m super_gradients.evaluate_checkpoint --experiment_namemy_experiment_name --ckpt_nameaverage_model.pth其中ckpt_name可传入ckpt_latest.pth、ckpt_best.pth、average_model.pth等任意检查点文件名缺省为ckpt_latest.pthckpt_root_dir缺省时使用仓库默认 checkpoints 目录。小结SuperGradients 的检查点体系围绕多时机保存、多模式加载、多途径恢复三个维度设计训练中自动产出ckpt_best.pth、ckpt_latest.pth、average_model.pth与ckpt_epoch_{N}.pth四类文件并按ckpt_root_dir/experiment_name/run_dir三级目录隔离加载时通过StrictLoad枚举在原生OFF/ON之外提供了NO_KEY_MATCHING按形状匹配与KEY_MATCHING按共有键匹配两种宽松策略配合load_backbone、load_ema_as_net、checkpoint_num_classes等能力覆盖迁移学习、换头微调等实战场景恢复训练则支持本地最新断点resume、指定 runrun_id、任意检查点分支resume_path、原始 recipe 自动恢复resume_experiment以及 WandB 远程续训五种途径。理解这套机制后无论是模型微调、实验分支管理还是生产环境断点续训都能在 SG 中高效落地。【免费下载链接】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),仅供参考