ARTICLE DETAIL

资讯详情

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

使用 SuperGradients 训练 CIFAR10 分类模型并完成迁移学习:完整实战指南

使用 SuperGradients 训练 CIFAR10 分类模型并完成迁移学习:完整实战指南 使用 SuperGradients 训练 CIFAR10 分类模型并完成迁移学习完整实战指南【免费下载链接】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 开源训练库以 ResNet18 在 CIFAR10 数据集上的完整训练流程为主线系统讲解从 Trainer 初始化、数据加载、模型构建、训练参数配置到断点续训、TensorBoard 可视化和 ImageNet 预训练权重迁移学习的全链路用法。读完本文你将掌握用几行代码搭建并运行一个图像分类训练任务、灵活覆盖默认配置以及在少量数据场景下借助预训练权重快速提升精度的实战方法。快速安装本示例唯一必需的 Python 包是super-gradients。安装它会自动带齐运行示例所需的全部依赖PyTorch、torchvision、hydra-core、PyYAML 等详见仓库根目录的 requirements.txt 与 setup.pypip install super-gradients1. 实验环境搭建初始化 TrainerTrainer是 SuperGradients 的训练入口它统一负责模型训练train验证集/测试集评估evaluate推理预测predict检查点checkpoint的保存与管理初始化Trainer只需两个关键参数experiment_name当前实验的唯一标识符将直接决定检查点与日志的落盘目录名ckpt_root_dir检查点、日志和 TensorBoard 文件的根目录。该参数可选若省略SuperGradients 默认使用项目根目录下的checkpoints目录。from super_gradients import Trainer experiment_name resnet18_cifar10_example CHECKPOINT_DIR /path/to/checkpoints/root/dir trainer Trainer(experiment_nameexperiment_name, ckpt_root_dirCHECKPOINT_DIR)仓库根目录下的checkpoints/目录正是这一默认约定的体现你可以把CHECKPOINT_DIR指向任何你拥有写权限的路径。2. 理解检查点目录结构检查点对渐进式训练、问题调试和模型部署至关重要。SuperGradients 会把它们组织成如下层级结构位于你指定的ckpt_root_dir之下ckpt_root_dir │ ├── experiment_name │ │ │ ├─── run_dir │ │ ├─ ckpt_best.pth # 验证指标最优时的权重 │ │ ├─ ckpt_latest.pth # 最近一个 epoch 结束时的权重 │ │ ├─ average_model.pth # 按指定 epoch 平均后的权重 │ │ ├─ ckpt_epoch_*.pth # 指定 epoch 的检查点如 epoch 10、15 等 │ │ ├─ events.out.tfevents.* # TensorBoard 运行产物 │ │ └─ log_timestamp.txt # 本次运行的 Trainer 日志 │ │ │ └─── other_run_dir │ └─ ... │ └─── other_experiment_name │ ├─── run_dir │ └─ ... │ └─── another_run_dir └─ ...各文件语义如下ckpt_best.pth当指定验证指标取得提升时保存即最佳模型ckpt_latest.pth每个 epoch 结束时更新即最新模型也是断点续训的默认加载对象average_model.pth当训练参数average_best_modelsTrue时生成是若干最优检查点权重的平均结果。关于检查点格式与加载规则的更多细节可参考专门的 Checkpoints.md 文档。在源码层面ckpt_best.pth与ckpt_latest.pth的文件名分别对应训练参数ckpt_best_name与ckpt_name默认值见 default_train_params.yaml。3. 数据集与 DataLoader本示例使用 CIFAR10 图像分类数据集。SuperGradients 内置了一批标准数据集与配套 DataLoader可直接开箱即用需要下载的数据集会自动下载会优雅地处理 DataLoader 创建并自带针对该数据集与模型架构调好的默认训练配方recipe。注意Trainer 与标准 PyTorchDataLoader/Dataset完全兼容。虽然本示例不展开但自定义 Dataset 和 DataLoader 完全可以无缝接入。3.A 使用 SuperGradients 默认 DataLoader创建训练与验证 DataLoader 只需两行代码from super_gradients.training import dataloaders train_dataloader dataloaders.get(namecifar10_train, dataset_params{}, dataloader_params{num_workers: 2}) valid_dataloader dataloaders.get(namecifar10_val, dataset_params{}, dataloader_params{num_workers: 2})这里调用了get()函数两次分别获取训练与验证 DataLoader。其参数含义为name字符串指定预置 DataLoader 的名称。SuperGradients 提供了大量预置 DataLoaderCIFAR10/100、ImageNet、COCO、Cityscapes 等本示例使用 CIFAR10 训练/验证集dataset_params字典用于覆盖训练配方recipe中定义的数据集相关默认参数。后续会演示如何用它替换图像预处理 transformsdataloader_params字典用于覆盖配方中定义的DataLoader 相关参数。这里示例性地把num_workers设为 2dataset一个torch.utils.data.Dataset对象用于接入自定义数据集实现不能与name或dataset_params同时传入。从源码看get()在传入自定义dataset时还会自动处理 sampler 参数与 collate_fn 参数_process_sampler_params、_process_collate_fn_params并把最终参数挂载到dataloader.dataloader_params上而在按name查找时会从注册表ALL_DATALOADERS中取出对应工厂函数如cifar10_train、cifar10_val见 dataloaders.py再基于cifar10_dataset_params配方实例化。完整实现见 dataloaders.py。我们可以随时打印 DataLoader 及其数据集的参数值import pprint print(Dataloader parameters:) pprint.pprint(train_dataloader.dataloader_params) print(Dataset parameters:) pprint.pprint(train_dataloader.dataset.dataset_params)预期输出Dataloader parameters: { batch_size: 256, drop_last: False, num_workers: 2, pin_memory: True, shuffle: True } Dataset parameters: { download: True, root: ./data/cifar10, target_transform: None, train: True, transforms: [ {RandomCrop: {size: 32, padding: 4}}, RandomHorizontalFlip, ToTensor, {Normalize: {mean: [0.4914, 0.4822, 0.4465], std: [0.2023, 0.1994, 0.201]}}, ] }这些默认值正是来自配方文件 cifar10_dataset_params.yaml训练集默认batch_size256、shuffleTruetransforms 依次为RandomCrop32×32padding 4、RandomHorizontalFlip、ToTensor与按 CIFAR10 数据集统计均值/方差做的Normalize验证集则额外先做Resize到 32×32。按上述方式调用get()时SuperGradients 会自动尝试下载 CIFAR10 数据集你会看到类似如下的下载日志定义好 DataLoader 后可以迭代它取出一批图像与标签用于可视化、检查张量形状等。例如可视化from matplotlib import pyplot as plt def show(images, labels, classes, rows6, columns5): fig plt.figure(figsize(10, 10)) for i in range(1, columns * rows 1): fig.add_subplot(rows, columns, i) plt.imshow(images[i-1].permute(1, 2, 0).clamp(0, 1)) plt.xticks([]) plt.yticks([]) plt.title(f{classes[labels[i-1]]}) plt.show() images_train, labels_train next(iter(train_dataloader)) show(images_train, labels_train, classestrain_dataloader.dataset.classes)输出可以看到图像经过了归一化。归一化过程是 SuperGradients 为 CIFAR10 数据集准备的默认训练配方的一部分。如后续小节所示SuperGradients 可以轻松覆盖全部或部分数据集与 DataLoader 参数从而在灵活性与开箱即用之间自由权衡。最后打印张量形状确认 batch 维度与通道/尺寸是否符合预期print(fTraining image tensor shape: {images_train.shape}) print(fTraining labels tensor shape: {labels_train.shape})输出Training image tensor shape: torch.Size([256, 3, 32, 32]) Training labels tensor shape: torch.Size([256])可以看到训练 DataLoader 的默认 batch size 为 256图像为 3 通道 32×32。3.B 覆盖数据集与 DataLoader 参数为展示 SuperGradients 在定制各训练组件上的灵活性我们覆盖数据集使用的 transforms 列表并使用torchvision的 transforms 来定义变换——这也体现了 SuperGradients 与 PyTorch 生态组件的无缝集成。为了可视化效果更直观这里只应用ToTensor()仅把输入图像转换为 PyTorch 张量from torchvision import transforms as T transforms_list [T.ToTensor()] vis_dataloader dataloaders.get(cifar10_train, dataset_params{transforms: transforms_list}, dataloader_params{num_workers: 2}) images, labels next(iter(vis_dataloader)) show(images, labels, classestrain_dataloader.dataset.classes)注意与默认 DataLoader 的唯一区别这里通过dataset_params字典传入了要覆盖的参数。运行结果如下替换 transforms 的效果一目了然——去掉归一化后图像中的物体细节更清晰可见。从实现上看dataset_params会以深度覆盖的方式合并进配方默认值后再传给数据集构造函数见 dataloaders.py 中dataset_params参数的说明。4. 模型架构定义本示例使用 ResNet18 架构。SuperGradients 内置了大量开箱即用的分类架构实现一行代码即可按所选架构定义模型from super_gradients.training import models from super_gradients.common.object_names import Models model models.get(model_nameModels.RESNET18, num_classes10)与获取预置 DataLoader 类似这里使用super_gradients.training.models的get()函数。上面传入两个参数model_name字符串指定架构名称取自 SuperGradients 提供的架构列表。Models.RESNET18对应的注册入口位于 resnet.pynum_classes整数模型需要预测的类别数会直接影响架构结构分类头输出维度。get()还支持一些常用参数完整签名见 model_factory.pyarch_params字典用于覆盖默认架构参数如残差块数量等checkpoint_path字符串外部检查点路径绝对或相对路径均可提供后会自动尝试加载该检查点pretrained_weights字符串指定模型预训练所用数据集的名称用于微调与迁移学习。pretrained_weights与checkpoint_path互斥同时传入会直接报错load_backbone仅加载检查点到model.backbone而非整个模型checkpoint_num_classes当检查点/预训练权重的类别数与num_classes不一致时使用——模型会先按检查点类别数初始化并加载权重再调用replace_head(new_num_classesnum_classes)替换分类头这正是从外部检查点做迁移学习的关键机制见 model_factory.pynum_input_channels输入通道数默认取模型默认值通常为 3。更多参数请参考该函数的 docstring。如前所述SuperGradients 与 PyTorch 高度兼容你也可以直接传入自定义的torch.nn.Module作为模型架构以获得最大灵活性本示例不展开。5. 训练配置到目前为止我们已经定义了 Trainer、数据集、DataLoader 和模型架构。开始训练前还需要定义训练参数。与前面一样SuperGradients 为本场景提供了调优好的默认训练参数一行代码即可获取from super_gradients.training import training_hyperparams training_params training_hyperparams.get(config_nametraining_hyperparams/cifar10_resnet_train_params)代码风格高度一致——获取训练参数同样调用get()函数它接受两个参数config_name字符串指定 recipes 目录中的 .yaml 配置文件名overriding_params可选字典用于覆盖加载到的训练参数。从实现看training_hyperparams.get()通过load_recipe加载完整配方、经 hydra 实例化后取出training_hyperparams字段再用override_default_params_without_nones合并覆盖参数实现见 training_hyperparams.py。本示例使用的 cifar10_resnet_train_params.yaml 通过defaults继承 default_train_params.yaml并定义了max_epochs250、StepLR 分段衰减epoch 100/150/200 处 lr×0.1、CrossEntropyLoss、SGD 优化器momentum 0.9、weight_decay 1e-4、以Accuracy为监控指标等默认值。打印训练参数即可看到全部可选项pprint.pprint(Training parameters:) pprint.pprint(training_params)输出Training parameters{ _convert_: all, average_best_models: True, batch_accumulate: 1, ckpt_best_name: ckpt_best.pth, ckpt_name: ckpt_latest.pth, clip_grad_norm: None, cosine_final_lr_ratio: 0.01, criterion_params: {}, dataset_statistics: False, ema: False, ema_params: {decay: 0.9999, decay_type: exp, beta: 15}, enable_qat: False, greater_metric_to_watch_is_better: True, initial_lr: 0.1, launch_tensorboard: False, load_opt_params: True, log_installed_packages: True, loss: CrossEntropyLoss, lr_cooldown_epochs: 0, lr_decay_factor: 0.1, lr_mode: StepLRScheduler, lr_schedule_function: None, lr_updates: array([100, 150, 200]), lr_warmup_epochs: 0, lr_warmup_steps: 0, max_epochs: 250, max_train_batches: None, max_valid_batches: None, metric_to_watch: Accuracy, mixed_precision: False, optimizer: SGD, optimizer_params: {weight_decay: 0.0001, momentum: 0.9}, phase_callbacks: [], pre_prediction_callback: None, precise_bn: False, precise_bn_batch_size: None, qat_params: {start_epoch: 0, quant_modules_calib_method: percentile, per_channel_quant_modules: False, calibrate: True, calibrated_model_path: None, calib_data_loader: None, num_calib_batches: 2, percentile: 99.99}, resume: False, resume_path: None, run_validation_freq: 1, save_ckpt_epoch_list: [], save_model: True, save_tensorboard_to_s3: False, seed: 42, sg_logger: base_sg_logger, sg_logger_params: {tb_files_user_prompt: False, launch_tensorboard: False, tensorboard_port: None, save_checkpoints_remote: False, save_tensorboard_remote: False, save_logs_remote: False, monitor_system: True}, silent_mode: False, step_lr_update_freq: None, sync_bn: False, tb_files_user_prompt: False, tensorboard_port: None, train_metrics_list: [Accuracy, Top5], valid_metrics_list: [Accuracy, Top5], warmup_initial_lr: None, warmup_mode: LinearEpochLRWarmup, zero_weight_decay_on_bias_and_bn: False }从输出可以看到大量可调项。这里挑选几个在 default_train_params.yaml 中有明确注释的要点说明学习率策略lr_mode支持StepLRScheduler/PolyLRScheduler/CosineLRScheduler/ExponentialLRScheduler/FunctionLRScheduler等lr_updates为阶梯下降的 epoch 位置lr_decay_factor为衰减倍率lr_warmup_epochs控制预热轮数优化器与损失optimizer可为Adam/SGD/RMSProploss可为 SuperGradients 内置损失或任意torch.nn.Module并配合criterion_params传参精度与效率mixed_precision开启混合精度emaTrue启用指数滑动平均ema_params中decay0.9999、decay_typeexpbatch_accumulate控制梯度累积步数验证与保存run_validation_freq控制验证频率save_ckpt_epoch_list可指定额外保存的 epochseed42保证可复现性调试辅助max_train_batches/max_valid_batches非 None 时会在迭代到该 batch 数后提前跳出循环方便快速调试。也可以在获取训练参数后直接修改例如training_params[max_epochs] 15 training_params[sg_logger_params][launch_tensorboard] True6. 训练、断点续训与迁移学习6.A 启动训练一切就绪后把模型、训练/验证 DataLoader 和训练参数接入 Trainer 的train()函数即可trainer.train(modelmodel, training_paramstraining_params, train_loadertrain_dataloader, valid_loadervalid_dataloader)训练进度会实时打印到屏幕[2023-02-01 20:57:27] INFO - sg_trainer_utils.py - TRAINING PARAMETERS: - Mode: Single GPU - Number of GPUs: 1 (4 available on the machine) - Dataset size: 50000 (len(train_set)) - Batch size per GPU: 256 (batch_size) - Batch Accumulate: 1 (batch_accumulate) - Total batch size: 256 (num_gpus * batch_size) - Effective Batch size: 256 (num_gpus * batch_size * batch_accumulate) - Iterations per epoch: 195 (len(train_set) / total_batch_size) - Gradient updates per epoch: 195 (len(train_set) / effective_batch_size) [2023-02-01 20:57:27] INFO - sg_trainer.py - Started training for 15 epochs (0/14) Train epoch 0: 100%|██████████| 196/196 [00:1800:00, 10.51it/s, Accuracy0.262, CrossEntropyLoss2.37, Top50.787, gpu_mem0.371] Validation epoch 0: 100%|██████████| 20/20 [00:0300:00, 6.12it/s] SUMMARY OF EPOCH 0 ├── Training │ ├── Accuracy 0.262 │ ├── CrossEntropyLoss 2.3702 │ └── Top5 0.787 └── Validation ├── Accuracy 0.3459 ├── CrossEntropyLoss 1.8811 └── Top5 0.871 训练开始时会打印训练参数摘要训练模式CPU/单 GPU/分布式训练、GPU 数量、训练集大小、batch size、每个 epoch 的迭代次数等。每个 epoch 的训练与验证进度条会展示配方中定义的跟踪指标Accuracy、loss、Top5 与 GPU 显存占用。每个 epoch 结束时打印训练/验证指标汇总后续 epoch 还会与历史最优值Best until now和上一 epochEpoch N-1对比 SUMMARY OF EPOCH 15 ├── Training │ ├── Accuracy 0.7594 │ │ ├── Best until now 0.7458 (↗ 0.0136) │ │ └── Epoch N-1 0.7458 (↗ 0.0136) │ ├── CrossEntropyLoss 0.686 │ │ ├── Best until now 0.7187 (↘ -0.0327) │ │ └── Epoch N-1 0.7187 (↘ -0.0327) │ └── Top5 0.9867 │ ├── Best until now 0.9849 (↗ 0.0019) │ └── Epoch N-1 0.9849 (↗ 0.0019) └── Validation ├── Accuracy 0.7425 │ ├── Best until now 0.746 (↘ -0.0035) │ └── Epoch N-1 0.7306 (↗ 0.0119) ├── CrossEntropyLoss 0.7331 │ ├── Best until now 0.7315 (↗ 0.0016) │ └── Epoch N-1 0.8048 (↘ -0.0717) └── Top5 0.9831 ├── Best until now 0.9838 (↘ -0.0007) └── Epoch N-1 0.9818 (↗ 0.0013) 每个 epoch 结束后各类日志与检查点都会保存到ckpt_root_dir与experiment_name决定的路径下。上图中第 15 个 epoch 的验证精度约 73%并不算高。为了解训练状态我们查看 TensorBoard 日志。6.B 查看 TensorBoard 日志在实验目录下打开终端执行tensorboard --logdir.也可以在任何位置用实验的完整路径运行该命令。SuperGradients 会把大量有用指标写入 TensorBoard包括 CPU/GPU 占用、学习率调度曲线、训练与验证 loss 及其他指标等。下面查看训练与验证 loss从图中可以看出训练及验证loss 在训练结束前尚未收敛说明继续训练更多 epoch 很可能进一步提升性能。前面修改训练参数时我们把max_epochs设成了 15下面让模型再续训 10 个 epoch。6.C 从检查点继续训练续训利用models.get()的checkpoint_path参数加载检查点。本例想从最近一次训练的终点继续因此加载ckpt_latest.pth。同时需要让 Trainer 知道这是继续训练而非从头开始——把训练参数resume设为True。最后设置新的max_epochs再次调用train()import os model models.get(model_nameModels.RESNET18, num_classes10, checkpoint_pathos.path.join(CHECKPOINT_DIR, experiment_name, ckpt_latest.pth)) training_params[resume] True training_params[max_epochs] 25 trainer.train(modelmodel, training_paramstraining_params, train_loadertrain_dataloader, valid_loadervalid_dataloader)可以看到训练从 epoch 15 继续又训练了 10 个 epoch[2023-02-01 21:21:16] INFO - sg_trainer_utils.py - TRAINING PARAMETERS: - Mode: Single GPU - Number of GPUs: 1 (4 available on the machine) - Dataset size: 50000 (len(train_set)) - Batch size per GPU: 256 (batch_size) - Batch Accumulate: 1 (batch_accumulate) - Total batch size: 256 (num_gpus * batch_size) - Effective Batch size: 256 (num_gpus * batch_size * batch_accumulate) - Iterations per epoch: 195 (len(train_set) / total_batch_size) - Gradient updates per epoch: 195 (len(train_set) / effective_batch_size) [2023-02-01 21:21:16] INFO - sg_trainer.py - Started training for 10 epochs (15/24) Train epoch 15: 100%|██████████| 196/196 [00:1800:00, 10.52it/s, Accuracy0.764, CrossEntropyLoss0.668, Top50.987, gpu_mem0.422] Validation epoch 15: 100%|██████████| 20/20 [00:0300:00, 6.10it/s] SUMMARY OF EPOCH 15 ├── Training │ ├── Accuracy 0.7644 │ ├── CrossEntropyLoss 0.6684 │ └── Top5 0.9865 └── Validation ├── Accuracy 0.7539 ├── CrossEntropyLoss 0.7271 └── Top5 0.9841 最终模型在完成 25 个 epoch 后停止训练 SUMMARY OF EPOCH 25 ├── Training │ ├── Accuracy 0.8177 │ │ ├── Best until now 0.8147 (↗ 0.003) │ │ └── Epoch N-1 0.8147 (↗ 0.003) │ ├── CrossEntropyLoss 0.5211 │ │ ├── Best until now 0.5281 (↘ -0.007) │ │ └── Epoch N-1 0.5281 (↘ -0.007) │ └── Top5 0.9921 │ ├── Best until now 0.9919 (↗ 0.0002) │ └── Epoch N-1 0.9919 (↗ 0.0002) └── Validation ├── Accuracy 0.8201 │ ├── Best until now 0.7873 (↗ 0.0328) │ └── Epoch N-1 0.7534 (↗ 0.0667) ├── CrossEntropyLoss 0.525 │ ├── Best until now 0.6145 (↘ -0.0895) │ └── Epoch N-1 0.7517 (↘ -0.2266) └── Top5 0.9907 ├── Best until now 0.9883 (↗ 0.0024) └── Epoch N-1 0.983 (↗ 0.0077) 验证精度提升到了 82%效果明显改善。需要说明的是resumeTrue时默认从ckpt_name即ckpt_latest.pth续训你也可以通过resume_path显式指定要恢复的检查点文件或设置load_opt_params决定是否同时恢复优化器状态相关参数定义见 default_train_params.yaml。6.D 迁移学习Transfer Learning到目前为止我们都是从零训练即模型权重随机初始化。对于简单任务和大型数据集这通常够用但在数据量不足的场景下我们往往希望利用其他来源的知识——这就是迁移学习。本示例展示其最简单的形式用预训练权重初始化模型进行微调fine-tuning具体使用在 ImageNet 上预训练的权重。SuperGradients 为不同模型提供了多种开箱即用的预训练权重。只需对现有代码做一处小改动model models.get(model_nameModels.RESNET18, num_classes10, pretrained_weightsimagenet)这里给models.get()传入了pretrained_weights参数值为预训练权重对应数据集的名称。注意该参数与checkpoint_path互斥。从实现看当num_classes与预训练权重对应的类别数不一致时模型会先按预训练类别数实例化并加载权重再通过replace_head(new_num_classesnum_classes)替换分类头以适配 CIFAR10 的 10 类输出见 model_factory.py。其余训练流程与前面完全一致。为与从零训练对比同样训练 25 个 epoch最终结果 SUMMARY OF EPOCH 25 ├── Training │ ├── Accuracy 0.8242 │ │ ├── Best until now 0.8267 (↘ -0.0025) │ │ └── Epoch N-1 0.8267 (↘ -0.0025) │ ├── CrossEntropyLoss 0.5035 │ │ ├── Best until now 0.4998 (↗ 0.0037) │ │ └── Epoch N-1 0.4998 (↗ 0.0037) │ └── Top5 0.9924 │ ├── Best until now 0.9917 (↗ 0.0007) │ └── Epoch N-1 0.9917 (↗ 0.0007) └── Validation ├── Accuracy 0.8377 │ ├── Best until now 0.8062 (↗ 0.0315) │ └── Epoch N-1 0.806 (↗ 0.0317) ├── CrossEntropyLoss 0.4834 │ ├── Best until now 0.5731 (↘ -0.0897) │ └── Epoch N-1 0.5785 (↘ -0.0952) └── Top5 0.9924 ├── Best until now 0.9903 (↗ 0.0021) └── Epoch N-1 0.9903 (↗ 0.0021) 相比随机初始化模型验证精度提升了约 1.7 个百分点82.01% → 83.77%。若想借助预训练权重获得更大提升通常还需要针对微调场景仔细调整学习率等超参数例如使用更小的initial_lr并配合余弦退火调度相关配置可参考 imagenet_resnet50.yaml 等配方中的做法。7. 用训练好的模型做预测现在我们有了一个性能合理的训练模型可以对新数据进行推理。首先导入所需包from PIL import Image import torch import numpy as np import requests接着加载训练好的权重注意这里加载的是ckpt_best.pth最佳检查点并切换为评估模式model models.get(model_nameModels.RESNET18, num_classes10, checkpoint_pathos.path.join(CHECKPOINT_DIR, experiment_name, ckpt_best.pth)) model.eval()我们想用一张属于训练类别之一的图像来测试模型例如一张青蛙照片。加载的图像必须经过与训练图像相同的变换模型才能正常工作url https://www.aquariumofpacific.org/images/exhibits/Magnificent_Tree_Frog_900.jpg image np.array(Image.open(requests.get(url, streamTrue).raw)) transforms T.Compose([ T.ToTensor(), T.Normalize(mean(0.4914, 0.4822, 0.4465), std(0.2023, 0.1994, 0.2010)), T.Resize((32, 32)) ]) input_tensor transforms(image).unsqueeze(0).to(next(model.parameters()).device)注意这里使用的归一化均值/方差与训练配方中 CIFAR10 的统计值完全一致见 cifar10_dataset_params.yaml且用Resize((32, 32))把任意尺寸的输入图缩放到训练时的 32×32。然后只需一行代码即可得到模型预测predictions model(input_tensor)查看模型预测结果plt.xlabel(train_dataloader.dataset.classes[torch.argmax(predictions)]) plt.imshow(image) plt.show()可以看到模型正确预测出输入图像是一只青蛙。8. 完整代码以下脚本整合了训练、从检查点续训与模型预测的完整流程只需修改CHECKPOINT_DIR变量即可直接运行from super_gradients import Trainer from super_gradients.training import dataloaders from super_gradients.training import models from super_gradients.common.object_names import Models from super_gradients.training import training_hyperparams import os from torchvision import transforms as T from PIL import Image import torch import numpy as np import requests import matplotlib.pyplot as plt def run(experiment_name, CHECKPOINT_DIR): # INITIALIZE TRAINER trainer Trainer(experiment_nameexperiment_name, ckpt_root_dirCHECKPOINT_DIR) # INITIALIZE DATALOADERS train_dataloader dataloaders.get(namecifar10_train, dataset_params{}, dataloader_params{num_workers: 2}) valid_dataloader dataloaders.get(namecifar10_val, dataset_params{}, dataloader_params{num_workers: 2}) # DEFINE MODEL model models.get(model_nameModels.RESNET18, num_classes10) # DEFINE TRAINING PARAMETERS training_params training_hyperparams.get(config_nametraining_hyperparams/cifar10_resnet_train_params) training_params[max_epochs] 15 # TRAIN trainer.train(modelmodel, training_paramstraining_params, train_loadertrain_dataloader, valid_loadervalid_dataloader) # LOAD MODEL FROM CHECKPOINT model models.get(model_nameModels.RESNET18, num_classes10, checkpoint_pathos.path.join(CHECKPOINT_DIR, experiment_name, ckpt_latest.pth)) # RESUME TRAINING training_params[resume] True training_params[max_epochs] 25 trainer.train(modelmodel, training_paramstraining_params, train_loadertrain_dataloader, valid_loadervalid_dataloader) # LOAD BEST CHECKPOINT model models.get(model_nameModels.RESNET18, num_classes10, checkpoint_pathos.path.join(CHECKPOINT_DIR, experiment_name, ckpt_best.pth)) model.eval() # PREDICT CLASS FOR TEST IMAGE url https://www.aquariumofpacific.org/images/exhibits/Magnificent_Tree_Frog_900.jpg image np.array(Image.open(requests.get(url, streamTrue).raw)) transforms T.Compose([ T.ToTensor(), T.Normalize(mean(0.4914, 0.4822, 0.4465), std(0.2023, 0.1994, 0.2010)), T.Resize((32, 32)) ]) input_tensor transforms(image).unsqueeze(0).to(next(model.parameters()).device) predictions model(input_tensor) plt.xlabel(train_dataloader.dataset.classes[torch.argmax(predictions)]) plt.imshow(image) plt.show() if __name__ __main__: experiment_name resnet18_cifar10_example CHECKPOINT_DIR /path/to/checkpoints/root/dir run(experiment_name, CHECKPOINT_DIR)延伸阅读若想用 YAML 配方而非代码方式驱动训练可参考 train_from_recipe.py 与 Recipes_Training.md以及本文多次引用的配方文件cifar10_resnet.yaml、default_train_params.yaml、cifar10_dataset_params.yaml关于检查点机制的完整说明见 Checkpoints.md更多分类任务示例含外部自定义数据集接入可参考 Example_Training-an-external-model.md 与 train_from_recipe_with_dataset_registry 示例目录仓库单元测试中的 pretrained_models_unit_test.py 与 test_models_factory.py 覆盖了models.get()的各类加载路径可作为理解模型工厂行为的参考。【免费下载链接】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),仅供参考
返回列表