MXNet Gluon Fit API 实战指南:用两行代码完成深度学习模型训练

发布时间:2026/9/21 1:33:41
MXNet Gluon Fit API 实战指南:用两行代码完成深度学习模型训练 深度学习机器学习人工智能【免费下载链接】mxnetLightweight, Portable, Flexible Distributed/Mobile Deep Learning with Dynamic, Mutation-aware Dataflow Dep Scheduler; for Python, R, Julia, Scala, Go, Javascript and more项目地址https://gitcode.com/gh_mirrors/mxnet1/mxnet点击查看免费下载本文基于 Apache MXNet 官方教程编写系统讲解 Gluoncontrib.estimator.EstimatorFit API的完整用法从 Fashion-MNIST 数据准备、ResNet-18 模型构建到使用 Fit API 以极简代码完成训练、验证、Checkpoint 保存与自定义事件处理器Event Handler。读完本文你将掌握 Fit API 的基本用法与高级定制能力理解Estimator内部训练循环与默认事件处理器的实现原理并能独立用更少代码训练自己的 Gluon 模型。一、Fit API 是什么把训练循环封装成一行fit()在传统 Gluon 训练流程中你需要手动编写训练循环training loop逐 batch 取数据、autograd.record()记录前向、计算 loss、反向传播、trainer.step()更新参数还要处理每个 epoch 结束时的指标统计与日志输出。这些样板代码boilerplate code占据了训练脚本的很大篇幅。Gluon Fit API即mxnet.gluon.contrib.estimator中的Estimator类正是为解决这一问题而设计你只需指定网络net、损失函数loss和训练数据即可开始训练。Fit API 会自动完成数据分批遍历、前向/反向、参数更新、指标计算与日志记录等所有琐碎环节高级用户仍可通过事件处理器Event Handler对训练阶段做细粒度控制甚至自行实现 bespoke 训练循环。从源码结构看Fit API 由三个文件组成estimator.py、event_handler.py、utils.pyEstimator训练与验证流程的调度器负责参数初始化、默认事件处理器的组装与训练循环的驱动事件处理器定义train_begin、train_end、batch_begin、batch_end、epoch_begin、epoch_end六个回调阶段用于指标计算、验证、日志、断点保存与早停工具函数负责指标合法性校验以及为特定损失函数推荐默认指标。对应的单元测试位于 test_gluon_estimator.py 与 test_gluon_event_handler.py覆盖了fit、验证、初始化器、trainer、指标、上下文与默认处理器等核心路径。二、环境准备与前置条件运行本教程需要MXNet版本 1.5.0可通过pip install mxnet安装 1.5.0 版本的 pip 包或直接从 master 分支源码编译Jupyter Notebook用于交互式运行教程对应的.ipynb文件。导入所需的包import mxnet as mx from mxnet import gluon from mxnet.gluon.model_zoo import vision from mxnet.gluon.contrib.estimator import estimator from mxnet.gluon.contrib.estimator.event_handler import TrainBegin, TrainEnd, EpochEnd, CheckpointHandler gpu_count mx.context.num_gpus() ctx [mx.gpu(i) for i in range(gpu_count)] if gpu_count 0 else mx.cpu()上面这段代码自动检测当前机器的 GPU 数量有 GPU 时把所有 GPU 组成上下文列表实现多卡训练没有 GPU 则回退到 CPU。Estimator._check_context会对传入的 context 做校验见 estimator.py只接受Context或Context列表若你不传 context默认使用gpu(0)多卡时会给出提示无 GPU 时默认cpu()。三、数据集准备Fashion-MNIST 与数据变换本教程使用Fashion-MNIST数据集训练一个图像分类模型。该数据集包含十个类别的时尚物品t-shirt/topT 恤/上衣、trouser裤子、pullover套头衫、dress连衣裙、coat外套、sandal凉鞋、shirt衬衫、sneaker运动鞋、bag包和 ankle boot短靴。训练集60,000 张 28 × 28 的灰度图像测试/验证集10,000 张 28 × 28 的灰度图像。使用gluon.data.vision包可以直接导入数据集并完成预处理# Get the training data fashion_mnist_train gluon.data.vision.FashionMNIST(trainTrue) # Get the validation data fashion_mnist_val gluon.data.vision.FashionMNIST(trainFalse)原始图像是 28 × 28 的灰度图而 ResNet-18 的输入尺寸是 224 × 224因此需要先Resize(224)再通过ToTensor()转为张量并归一化最后用Compose串联所有变换transforms [gluon.data.vision.transforms.Resize(224), # 模型输入尺寸为 224 gluon.data.vision.transforms.ToTensor()] # 将所有变换堆叠在一起 transforms gluon.data.vision.transforms.Compose(transforms)transform_first只对数据集的第一个元素即图像应用变换而标签保持不变# Apply the transformations fashion_mnist_train fashion_mnist_train.transform_first(transforms) fashion_mnist_val fashion_mnist_val.transform_first(transforms)接下来构造DataLoader。batch_size256表示每批图像数量num_workers4表示并行加载数据的 worker 数量。训练数据需要shuffleTrue打乱顺序以提升收敛效果验证数据无需打乱batch_size 256 # 图像批大小 num_workers 4 # 使用 DataLoader 加载数据时的并行 worker 数量 train_data_loader gluon.data.DataLoader(fashion_mnist_train, batch_sizebatch_size, shuffleTrue, num_workersnum_workers) val_data_loader gluon.data.DataLoader(fashion_mnist_val, batch_sizebatch_size, shuffleFalse, num_workersnum_workers)四、模型、损失函数与优化器从Gluon Model Zoo预训练模型与网络结构定义的集合加载 resnet-18 网络结构并使用 Xavier 初始化器初始化参数。这里我们只使用 model zoo 中的网络结构从零开始训练resnet_18_v1 vision.resnet18_v1(pretrainedFalse, classes 10) resnet_18_v1.initialize(init mx.init.Xavier(), ctxctx)由于是十类多分类问题损失函数选用SoftmaxCrossEntropyLoss优化器选用sgd随机梯度下降。你也可以按需更换其他损失函数或优化器。loss_fn gluon.loss.SoftmaxCrossEntropyLoss()然后创建 trainer 对象。gluon.Trainer接收网络的参数集合、优化器名称与超参learning_rate 0.04 # 可自行调整学习率 num_epochs 2 # 可增加训练轮数 trainer gluon.Trainer(resnet_18_v1.collect_params(), sgd, {learning_rate: learning_rate})值得说明的是Estimator的构造函数中trainer是可选参数。如果你不传 trainer从源码 estimator.py 可以看到它会默认使用学习率为 0.001 的 SGD 优化器并给出警告initializer同理——若网络尚未初始化会按你传入的初始化器缺省用默认初始化器初始化若网络已初始化则跳过并提示可用force_reinitTrue强制重初始化。五、Fit API 基本用法两行代码开始训练如前面所述Fit API 大幅简化了训练样板代码。基本用法只需两步train_acc mx.metric.Accuracy() # 用于监控的指标 # 定义 estimator传入模型、损失函数、指标、trainer 对象与上下文 est estimator.Estimator(netresnet_18_v1, lossloss_fn, metricstrain_acc, trainertrainer, contextctx) # 忽略 nightly 测试CI中的警告 import warnings with warnings.catch_warnings(): warnings.simplefilter(ignore) # 魔法行开始训练 est.fit(train_datatrain_data_loader, epochsnum_epochs)Estimator构造参数一览与 estimator.py 中的签名一致参数类型说明netgluon.Block用于训练的模型lossgluon.loss.Loss训练时计算的损失目标函数必须是gluon.loss.Loss实例metricsEvalMetric或EvalMetric列表模型评估指标initializerInitializer网络参数初始化器可选trainerTrainer对网络参数应用优化器的 trainer可选contextContext或Context列表训练运行的设备fit()的完整签名是fit(train_data, val_dataNone, epochsNone, event_handlersNone, batchesNone, batch_axis0)见 estimator.py几点重要约束train_data与val_data必须是gluon.data.DataLoader若你手头是DataIter或 NDArray需先转换为 DataLoaderepochs与batches必须且只能指定一个——要么按 epoch 数迭代要么按 batch 数迭代两者同时指定或都不指定会抛出ValueErrorbatch_axis0表示按第 0 维将 batch 切分到各设备多卡场景。训练开始后输出日志Training begin: using optimizer SGD with current learning rate 0.0400 Train for 2 epochs. [Epoch 0] finished in 25.110s: train_accuracy : 0.7877 train_softmaxcrossentropyloss0 : 0.5905 [Epoch 1] finished in 23.595s: train_accuracy : 0.8823 train_softmaxcrossentropyloss0 : 0.3197 Train finished using total 48s at epoch 1. train_accuracy : 0.8823 train_softmaxcrossentropyloss0 : 0.3197注意日志中除了我们传入的Accuracy指标外还自动出现了一个 loss 指标。这是因为Estimator._add_default_training_metrics见 estimator.py做了两件事当没有指定指标时根据损失函数类型推荐默认指标——从 utils.py 可以看到对于SoftmaxCrossEntropyLoss会推荐Accuracy()同时总是把 loss 包装成指标追加进训练指标列表用于记录每个 batch/epoch 的损失值。训练循环内部发生了什么fit()并非魔法它的执行流程在 estimator.py 中清晰可见依次调用所有 handler 的train_begin进入 epoch 循环调用epoch_begin遍历train_data的每个 batch先batch_begin然后调用fit_batch(batch, batch_axis)最后batch_end若任一 handler 在batch_end返回了停止信号如达到最大 batch 数提前退出本 epochepoch 结束后调用所有epoch_end若任一 handler 返回停止信号如达到最大 epoch 数结束训练最后调用所有train_end。其中单 batch 的训练逻辑fit_batch见 estimator.py等价于手写训练循环split_and_load将数据按设备切分 →autograd.record()记录前向 → 计算预测与 loss → 对每个 lossbackward()→trainer.step(batch_size)更新参数。也就是说Fit API 只是把这些步骤封装起来其底层机制与你手写的训练循环完全一致。六、高级用法事件处理器Event HandlerFit API 的高度可定制性来自Event Handler机制。它提供 6 个可覆写的回调方法覆盖训练的各阶段回调方法触发时机对应的基类train_begin训练开始TrainBegintrain_end训练结束TrainEndepoch_begin每个 epoch 开始EpochBeginepoch_end每个 epoch 结束EpochEndbatch_begin每个 batch 开始BatchBeginbatch_end每个 batch 结束BatchEnd六个基类定义在 event_handler.py。Estimator._categorize_handlers见 estimator.py会根据 handler 继承的基类把它们分类到 6 个回调列表中只调用实际实现了对应方法的 handler。6.1 内置事件处理器Fit API 提供了三类开箱即用的内置 handler都在mxnet.gluon.contrib.estimator.event_handler中导出LoggingHandler记录超参数、训练统计与过程信息可指定file_name/file_location将日志写入文件verbose控制粒度LOG_PER_EPOCH1每 epoch 打印一次LOG_PER_BATCH2每个 batch 打印一次CheckpointHandler按指定周期保存模型。首次 batch 结束时若网络是已 hybridize 的HybridBlock会额外保存网络结构-symbol.json之后按epoch_period默认每 epoch或batch_period保存.params参数与.states训练器状态EarlyStoppingHandler监控某个指标当指标在patience个 epoch 内未改善改善幅度小于min_delta时提前终止训练避免在性能停滞时浪费时间。此外Estimator默认会为每次fit装配 3 个工具型 handler见 estimator.pyStoppingHandler根据epochs或batches决定训练何时结束对应fit传入的迭代次数MetricHandler在每个 batch 结束时用预测、标签与 loss 更新训练指标并在每个 epoch 开始时reset()指标它的priority被设为负无穷保证先于其他 handler 执行其他 handler 才能拿到最新指标值ValidationHandler在每个 epoch 结束时用val_data调用Estimator.evaluate计算验证指标支持epoch_period与batch_period控制验证频率。默认 handler 的优先级设计值得注意MetricHandler和ValidationHandler的priority -np.Inf最先执行LoggingHandler的priority np.Inf最后执行handler 列表会按 priority 排序见 estimator.py。如果你自定义的 handler 依赖指标值或需要覆盖日志可参考这一约定设置自己的 priority。如果你自己传入同类 handler例如自定义的ValidationHandler或LoggingHandlerEstimator._prepare_default_handlers会跳过对应的默认 handler实现覆盖同时会校验所有 handler 引用的指标必须来自estimator.train_metrics或estimator.val_metrics校验逻辑见 utils.py确保指标引用一致。6.2 自定义事件处理器示例下面演示如何通过继承多个基类来编写一个自定义 handler。该 handler 的功能很简单在每个 epoch 结束时记录损失值。注意每个回调方法都会被传入Estimator对象因此你可以通过estimator访问训练指标。class LossRecordHandler(TrainBegin, TrainEnd, EpochEnd): def __init__(self): super(LossRecordHandler, self).__init__() self.loss_history {} def train_begin(self, estimator, *args, **kwargs): print(Training begin) def train_end(self, estimator, *args, **kwargs): # 训练结束时打印所有损失 print(Training ended) for loss_name in self.loss_history: for i, loss_val in enumerate(self.loss_history[loss_name]): print(Epoch: {}, Loss name: {}, Loss value: {}.format(i, loss_name, loss_val)) def epoch_end(self, estimator, *args, **kwargs): for metric in estimator.train_metrics: # 在训练指标中查找 loss # 我们已将 loss 值包装成指标用于记录 if isinstance(metric, mx.metric.Loss): loss_name, loss_val metric.get() # 记录该 epoch 的 loss 值 self.loss_history.setdefault(loss_name, []).append(loss_val)由于 loss 被包装成了mx.metric.Loss指标这正是前面提到的_add_default_training_metrics的行为我们只需遍历estimator.train_metrics、用isinstance(metric, mx.metric.Loss)识别 loss 指标再调用metric.get()即可拿到当前 epoch 的 loss 名称与数值。七、带验证与 Checkpoint 的完整训练流程将自定义 handler 与内置的CheckpointHandler组合使用即可同时完成验证指标统计、断点保存与 loss 记录。首先重置模型、trainer 与指标因为上面的实验已经训练过两轮# 重置上面的模型、trainer 与 accuracy 对象 resnet_18_v1.initialize(force_reinitTrue, init mx.init.Xavier(), ctxctx) trainer gluon.Trainer(resnet_18_v1.collect_params(), sgd, {learning_rate: learning_rate}) train_acc mx.metric.Accuracy()然后重新定义 estimator并传入两个事件处理器# 定义 estimator传入模型、损失函数、指标、trainer 对象与上下文 est estimator.Estimator(netresnet_18_v1, lossloss_fn, metricstrain_acc, trainertrainer, contextctx) # 定义内置的 CheckpointHandler checkpoint_handler CheckpointHandler(model_dir./, model_prefixmy_model, monitortrain_acc, # 监控某个指标 save_bestTrue) # 保存指标最优时的模型 # 实例化前面自定义的 handler loss_record_handler LossRecordHandler() # 忽略 nightly 测试CI中的警告 import warnings with warnings.catch_warnings(): warnings.simplefilter(ignore) # 魔法行 est.fit(train_datatrain_data_loader, val_dataval_data_loader, epochsnum_epochs, event_handlers[checkpoint_handler, loss_record_handler]) # 传入事件处理器训练输出同时包含验证指标与自定义 handler 的打印信息Training begin: using optimizer SGD with current learning rate 0.0400 Train for 2 epochs. [Epoch 0] finished in 25.236s: train_accuracy : 0.7917 train_softmaxcrossentropyloss0 : 0.5741 val_accuracy : 0.6612 val_softmaxcrossentropyloss0 : 0.8627 [Epoch 1] finished in 24.892s: train_accuracy : 0.8826 train_softmaxcrossentropyloss0 : 0.3229 val_accuracy : 0.8474 val_softmaxcrossentropyloss0 : 0.4262 Train finished using total 50s at epoch 1. train_accuracy : 0.8826 train_softmaxcrossentropyloss0 : 0.3229 val_accuracy : 0.8474 val_softmaxcrossentropyloss0 : 0.4262 Training begin Epoch 1, loss 0.5741 Epoch 2, loss 0.32297.1 CheckpointHandler 关键参数详解从源码 event_handler.py 可见CheckpointHandler支持以下参数参数默认值说明model_dir必填保存模型架构、参数与训练器状态文件的目录不存在会自动创建model_prefixmodel所有 checkpoint 文件名的前缀monitorNone用于判断模型是否改善的指标save_bestTrue时必填verbose0详细程度1表示每次保存 checkpoint 时打印信息save_bestFalse为True时按监控指标的最优值保存最佳模型文件名带-best后缀monitor不能为Nonemodeauto{auto, min, max}之一决定指标改善的比较方向auto模式下对名称含acc/f1的指标用greater否则用lessepoch_period1每隔多少 epoch 保存一次网络batch_periodNone每隔多少 batch 保存一次网络默认不按 batch 保存max_checkpoints5model_dir中最多保留的 checkpoint 文件数更旧的会被自动删除-best文件不计入resume_from_checkpointFalse是否从model_dir中的 checkpoint 恢复训练恢复后按剩余 epoch/batch 继续训练保存的文件分为三类见_save_params_and_trainer与_save_symbolevent_handler.py网络结构文件{model_prefix}-symbol.json——仅当模型是HybridBlock且已调用hybridize()时才会在第一个 batch 后保存参数文件{model_prefix}-epoch{epoch}batch{batch}.params或{model_prefix}-best.params通过net.save_parameters保存训练器状态文件同名.states文件通过trainer.save_states保存包含优化器动量等状态用于断点续训。7.2 加载已保存的模型使用 Gluon 的load_parametersAPI 即可加载保存的模型参数。更多细节可参考 加载模型参数教程resnet_18_v1 vision.resnet18_v1(pretrainedFalse, classes10) resnet_18_v1.load_parameters(./my_model-best.params, ctxctx)注意load_parameters需要网络结构一致这里重建了同样配置的resnet18_v1并把参数加载到指定ctx上。八、验证 Fit API 行为的测试用例如果你想深入了解 Fit API 的边界行为仓库中的单元测试是最直接的证据test_gluon_estimator.py 中的test_fit、test_validation、test_initializer、test_trainer、test_metric、test_loss、test_context分别验证了训练、验证流程、初始化器、自定义 trainer、指标与 loss 校验、上下文推断以及test_categorize_handlershandler 分类与test_default_handlers默认 handler 装配等内部逻辑test_gluon_event_handler.py 覆盖各事件处理器的行为test_estimator_cnn.py 提供了基于 CNN 的端到端训练示例。九、下一步学习Fit API 是快速上手的利器但深度学习训练远不止于此。进阶方向包括学习更多 Gluon 训练 API参考 Gluon 教程目录阅读本文引用的核心源码 estimator.py 与 event_handler.py理解fit_batch/evaluate_batch的可覆写设计——对于自定义训练逻辑可以继承Estimator并覆写这两个方法而无需重写整个训练循环结合本文介绍的CheckpointHandler(resume_from_checkpointTrue)与EarlyStoppingHandler为长时间训练任务构建断点续训 自动早停的完整方案。赞分享深度学习机器学习人工智能【免费下载链接】mxnetLightweight, Portable, Flexible Distributed/Mobile Deep Learning with Dynamic, Mutation-aware Dataflow Dep Scheduler; for Python, R, Julia, Scala, Go, Javascript and more项目地址https://gitcode.com/gh_mirrors/mxnet1/mxnet点击查看免费下载相关推荐img2vec的21个预训练模型全解析:ResNet、VGG、DenseNet、EfficientNet选型完全指南img2vec的21个预训练模型全解析:ResNet、VGG、DenseNet、EfficientNet选型完全指南 img2vec 是一个基于 PyTorch人工智能深度学习机器学习MXNet Gluon Fit API 实战用 Estimator 几行代码训练 ResNet-18 图像分类模型MXNet Gluon Fit API 实战用 Estimator 几行代码训练 ResNet 18 图像分类模型 Apache MXNet 的 Gluon深度学习人工智能机器学习分布式训练MXNet 速成课从 NDArray 到 GPU 多卡训练用 Gluon 走完深度学习全流程MXNet 速成课从 NDArray 到 GPU 多卡训练用 Gluon 走完深度学习全流程 本文是基于 Apache MXNet 官方 Crash Cou深度学习人工智能机器学习分布式训练上一篇如何在5分钟内快速上手PixiJS Live2D插件完整实战指南下一篇如何用Home Assistant实现榨汁机智能控制打造健康饮品自动化厨房 创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询