SuperGradients 模型检查点完全指南:保存、加载、恢复训练与评估
2026/9/18 22:32:35 网站建设 项目流程

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)是深度学习训练流程中"可以回退的锚点":它既记录了模型在训练各阶段的状态快照,也承载了中断后无缝恢复训练的能力。本文以 SuperGradients(SG)官方文档documentation/source/Checkpoints.md为主体,结合仓库内 Trainer 与 checkpoint 工具源码,系统讲解 SG 中检查点的保存时机与目录组织、内部数据结构、严格/宽松/按形状匹配等加载策略、基于 recipe 的恢复训练与远程(WandB)断点续训,以及一键评估历史检查点的方法。读完本文,你将能够熟练运用models.getload_checkpoint_to_modelTrainer.resume_experiment等 API,独立完成从保存到加载再到恢复、评估的完整闭环。

什么是模型检查点

训练过程中,模型性能会随着它看到的样本数量不断变化。业界最佳实践是在训练的关键节点保存模型状态——每一份保存下来的状态就是一个 checkpoint,对应模型开发过程中某个时刻的版本。训练结束后,应当使用验证集上表现最佳的 checkpoint 作为最终产物;同时,checkpoint 也保证了训练被中断时可以从断点继续,而不必从头再来。

SuperGradients 遵循这一思想,在训练全程自动保存多个不同用途的 checkpoint。每个 checkpoint 都对应一个训练阶段,彼此分工明确:有的用于追踪最佳性能,有的用于支持断点续训,有的用于产出最终部署模型。如果你还不熟悉 SG 的实验管理机制,建议先阅读 实验管理文档 了解ckpt_root_direxperiment_namerun等概念。

Checkpoint 保存策略:哪个文件、何时保存、存在哪里

四种默认检查点文件

在 SG 中,训练过程中会按照下表所示的时机自动保存不同类型的 checkpoint 文件:

Checkpoint 文件名保存时机
ckpt_best.pth每次验证时达到新的最佳metric_to_watch指标即覆盖保存
ckpt_latest.pth每个 epoch 结束时保存,持续覆盖为最新状态
average_model.pth训练结束时保存,由验证指标最优的 10 个模型快照平均得到;仅当average_best_models=True时生成
ckpt_epoch_{EPOCH_INDEX}.pthsave_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.pth);
  • save_model:总开关,控制是否保存模型检查点。

从源码看,这些逻辑集中在 sg_trainer.py 的_save_checkpoint中:每次验证结束后,先根据metric_to_watchgreater_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_direxperiment_name由用户在实例化Trainer时指定:
Trainer(ckpt_root_dir='path/to/ckpt_root_dir', experiment_name="my_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指标值(float)
epoch最近完成的一个 epoch
optimizer_state_dict优化器的state_dict
scaler_state_dict可选——仅当以mixed_precision=True训练时存在,为Trainer.scalerstate_dict
ema_net可选——仅当以ema=True训练时存在,为 EMA 模型的state_dict。注意:average_model.pth即使开启 EMA 也不含该键,因为平均模型本身就是对 EMA 快照做平均,其net键已经是 EMA 快照的平均结果
torch_scheduler_state_dict可选——仅当使用 PyTorch 原生 LR scheduler 时存在(见 LRScheduling)

对照源码,_save_checkpoint构建 state dict 时除上述键外还额外写入:

  • metrics:本次验证(及训练)所有指标的汇总 dict;
  • packages:训练环境已安装的包及版本列表,便于复现环境;
  • _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_remote=True,训练过程中产出的 checkpoint 就会被同步上传到远程存储。更完整的配置说明见 第三方实验监控文档。

远程保存的意义不仅在于备份,它还是"训练中断后从云端续训"的前提,具体操作见下文"从远程存储恢复训练"一节。

加载检查点:加载权重与恢复训练是两回事

加载检查点可以按使用场景拆分为两类:

  1. 仅加载模型权重(用于推理、迁移学习、微调);
  2. 完整恢复训练(连同优化器、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_name=Models.RESNET18, num_classes=10) train_params = { ... "loss": "CrossEntropyLoss", "criterion_params": {}, "save_ckpt_epoch_list": [10, 15] ... } trainer.train(model=model, training_params=train_params, train_loader=train_dataloader, valid_loader=valid_dataloader)

训练结束后,我们想加载ckpt_best.pth的权重,只需把它的完整路径传给models.getcheckpoint_path参数:

from super_gradients.training import models from super_gradients.common.object_names import Models model = models.get( model_name=Models.RESNET18, num_classes=10, 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_model

from 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_name=Models.RESNET18, num_classes=10) load_checkpoint_to_model( net=model, 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分别用strictstrict_load两个参数承担 PyTorchstrict的职责,但它们接受的是 SG 自定义的StrictLoad枚举类型。该枚举定义于 strict_load.py:

class StrictLoad(Enum): """ Wrapper for adding more functionality to torch's 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 SuperGradient's adapt_checkpoint function, which loads a checkpoint by matching each layer's 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(等价strict=False)与ON(等价strict=True),SG 还额外提供了两种宽松加载模式:

  • NO_KEY_MATCHING:利用state_dictOrderedDict这一事实,按层顺序做形状匹配加载,完全忽略键名匹配。当网络底层结构一致、但各层的state_dict键名与模型内键名不一致时非常有用。从源码看,adaptive_load_state_dict会先尝试按 strict 加载,失败后再调用adapt_state_dict_to_fit_model_layer_names将 checkpoint 键名重排为模型键名,最终以strict=True完成加载。
  • KEY_MATCHING:仅加载 checkpoint 与模型共有的键,且同一键下张量形状不一致的层会被跳过。其实现是 transfer_weights,逐个键尝试load_state_dict(..., strict=False),形状不兼容的层直接跳过。

下面用一个简单例子演示不同 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)。因此:

  • 使用strict=True(即StrictLoad.ON)从一个加载到另一个会直接报错崩溃;
  • 使用strict=False(即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_weights="coco")

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 checkpoint's 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_classes=num_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_load=cfg.checkpoint_params.strict_load, pretrained_weights=cfg.checkpoint_params.pretrained_weights, checkpoint_path=cfg.checkpoint_params.checkpoint_path, load_backbone=cfg.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_classes=num_classes)替换输出头,避免因 head 维度不匹配而加载失败;
  • external_checkpoint_pathsource_ckpt_folder_name用于加载不在默认 checkpoints 目录下的外部权重。

恢复训练(Resume Training)

SG 的断点续训由三个训练参数协同控制,它们都在 default_train_params.yaml 中有定义,提供了从"最新断点"到"任意指定检查点分支"的灵活度:

resume: False # 是否从同一实验名下的最新一次运行继续训练 run_id: # 同一实验内,要从哪一次运行(run)恢复 resume_path: # 直接指定一个 .pth 检查点文件的路径来恢复训练

1. 恢复最近一次运行

resume=True设置为真,SG 会在同一实验名下找到最近一次运行的ckpt_latest.pth(默认文件名由ckpt_name控制)并从该断点继续:

# 从 cifar_experiment 最近一次运行处继续 python -m super_gradients.train_from_recipe --config-name=cifar10_resnet experiment_name=cifar_experiment training_hyperparams.resume=True

2. 恢复指定运行

通过run_id可以精确恢复同一实验内的某次运行:

# 从 cifar_experiment 中由 run_id 标识的那次运行继续 python -m super_gradients.train_from_recipe --config-name=cifar10_resnet experiment_name=cifar_experiment run_id=RUN_20230802_131052_651906

3. 从指定检查点分支

通过resume_path指定任意.pth文件,SG 会新建一个 run 目录,从该检查点继续训练,并把新产生的检查点保存到新目录中——这非常适合做"分支实验":

# 从指定检查点分支,创建一次新的 run python -m super_gradients.train_from_recipe --config-name=cifar10_resnet experiment_name=cifar_experiment training_hyperparams.resume_path=/path/to/checkpoint.pth

从源码看,_load_checkpoint_to_model会综合resumerun_idresume_pathresume_from_remote_sg_logger四个来源决定是否加载检查点;其中resume=True且没有显式路径时沿用原 run,而使用resume_from_remote_sg_loggerresume_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,再注入resume=Truerun_id后重新走train_from_config)。

命令行等价用法可参考 resume_experiment.py 入口(旧示例脚本位于 examples/resume_experiment_example/resume_experiment.py,已标记弃用):

python -m super_gradients.resume_experiment --experiment_name=my_experiment_name

从远程存储恢复训练(WandB)

SG 支持从 SG Logger 定义的远程存储中恢复训练。前提是:训练期间已在sg_logger_params中开启save_checkpoints_remote=True,使检查点被同步到远程(例如 WandB run 的存储)。

假设我们使用 WandB SG Logger 运行实验,则training_hyperparams应包含:

sg_logger: wandb_sg_logger, # Weights&Biases 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, # W&B project name save_checkpoints_remote: True, save_tensorboard_remote: True, save_logs_remote: True, entity: <YOUR-ENTITY-NAME>, # username or team name where you're sending runs api_server: <OPTIONAL-WANDB-URL> # Optional: In case your experiment tracking is not hosted at wandb servers

save_checkpoints_remote=True会促使训练全程在 WandB 中保存检查点。若此时训练被中断,只需设置两个训练超参即可从 WandB run 存储中的检查点恢复:

  1. 设置resume_from_remote_sg_logger
resume_from_remote_sg_logger: True
  1. sg_logger_params中通过wandb_id传入原 run 的 id:
sg_logger: wandb_sg_logger, # Weights&Biases 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, # W&B project name save_checkpoints_remote: True, save_tensorboard_remote: True, save_logs_remote: True, entity: <YOUR-ENTITY-NAME>, # username or team name where you're 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_remote=True运行的实验有效。

评估检查点

与"用原始 recipe 恢复训练"的思路类似,我们常常希望在不重新熟悉训练配置的情况下直接评估某个历史检查点。为此 SG 提供了两个 Trainer 方法:

  • Trainer.evaluate_checkpoint(...):评估你自己此前某个实验产生的检查点,使用该实验训练时完全一致的参数(数据集、验证指标等)。即便之后 recipe 被修改过,评估仍按训练时的参数进行,确保验证结果与训练时完全可比。
  • Trainer.evaluate_recipe(...)(当前源码中已演进为evaluate_from_config):评估 Model Zoo 的预训练模型检查点,或以不同的参数评估检查点(例如更换数据集或验证指标)。

从源码看,evaluate_checkpoint的实现逻辑是:定位到指定实验(默认最近 run)的历史配置快照,注入resume=True与目标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_name=my_experiment_name --ckpt_name=average_model.pth

其中ckpt_name可传入ckpt_latest.pthckpt_best.pthaverage_model.pth等任意检查点文件名,缺省为ckpt_latest.pthckpt_root_dir缺省时使用仓库默认 checkpoints 目录。

小结

SuperGradients 的检查点体系围绕"多时机保存、多模式加载、多途径恢复"三个维度设计:训练中自动产出ckpt_best.pthckpt_latest.pthaverage_model.pthckpt_epoch_{N}.pth四类文件,并按ckpt_root_dir/experiment_name/run_dir三级目录隔离;加载时通过StrictLoad枚举在原生OFF/ON之外提供了NO_KEY_MATCHING(按形状匹配)与KEY_MATCHING(按共有键匹配)两种宽松策略,配合load_backboneload_ema_as_netcheckpoint_num_classes等能力覆盖迁移学习、换头微调等实战场景;恢复训练则支持本地最新断点(resume)、指定 run(run_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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询