news 2026/9/21 1:21:25

MXNet Gluon 迁移学习实战:从实验训练到模型部署的完整流程

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
MXNet Gluon 迁移学习实战:从实验训练到模型部署的完整流程

MXNet Gluon 迁移学习实战:从实验训练到模型部署的完整流程

【免费下载链接】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

导读

本文以「花种识别服务」这一典型业务场景为线索,完整讲解如何在 MXNet 中使用 Gluon API 走通「数据准备 → 迁移学习微调 → 模型序列化 → 部署推理」的全链路。你将掌握基于 ImageNet 预训练 ResNet50 V2 进行小数据集微调的方法、Gluon 数据增强与hybridize()的实战用法,以及如何把微调后的模型导出为.json+.params文件并用 Module API 完成在线推理,为后续接入 C++、Java、Scala 或 MXNet Model Server 部署打下基础。本文基于 docs/python_docs/python/tutorials/getting-started/gluon_from_experiment_to_deployment.md 展开,并补充了对应源码实现细节。

场景与思路:为什么用迁移学习

假设你要构建一个提供「花种识别」能力的服务。一个常见困境是:业务方往往没有足够的数据去从零训练一个高质量模型。此时可以借助迁移学习(Transfer Learning):利用一个在大规模标准数据集(如约 130 万张图像的 ImageNet)上预训练好的、解决相近任务的模型,把其中学到的通用视觉特征迁移到新任务上,从而用少量数据训练出更稳健的模型。

Gluon 模型动物园(Model Zoo)为分类、目标检测、语义分割等标准任务提供了众多预训练模型。本教程选用在 ImageNet 上预训练的ResNet50 V2(出自论文Identity Mappings in Deep Residual Networks),其 ImageNet top-1 精度为 77.11%。我们的目标是从中迁移尽可能多的知识,用于识别 Oxford 102 花卉数据集的 102 个花种。

说明:本教程的训练与推理部分使用 Python;完成本教程后,可继续阅读仓库中的 C++ 推理示例 了解如何在 C++ 侧加载同一套模型文件。

前置条件

  • 使用带 Python(Gluon)与 C++ 包的 MXNet 构建环境;
  • 具备 Gluon 基础知识(可先学习官方 Gluon 速成课所覆盖的 Block、Trainer、Dataset/DataLoader 等概念)。

数据准备:Oxford 102 花卉数据集

教程采用 Oxford 102 Category Flower Dataset(102 类花卉,共 8189 张图像)作为示例。仓库在 docs/tutorial_utils/data/oxford_102_flower_dataset.py 中提供了配套工具脚本,负责自动下载并整理数据。从源码看,该脚本主要完成三件事:

  1. download_data():从牛津官网下载102flowers.tgz(图像)、imagelabels.mat(标签)与setid.mat(训练/测试/验证集合划分),解压出jpg/目录;
  2. prepare_data():用scipy.io.loadmat读取三个.mat文件,将官方划分转换为 0-based 标签,并把每个样本映射为花名(如lotusrose),随后把jpg/中的图像按类别复制进train/test/valid/三个子目录;
  3. generate_synset():按字母序把 102 个花名写入synset.txt,供推理阶段映射标签序号到名称使用。

在仓库中运行如下代码即可下载并组织数据(工具脚本位于上述路径,可先将其复制到当前工作目录再导入):

import mxnet as mx data_util_file = "oxford_102_flower_dataset.py" mx.test_utils.download(base_url.format(data_util_file), fname=data_util_file) import oxford_102_flower_dataset # 下载并将数据整理到 train/test/valid 目录 path = './data' oxford_102_flower_dataset.get_data(path)

整理完成后,同一类别的图片会归入同一文件夹,目录结构与gluon.data.vision.ImageFolderDataset期望的root/类别/图片组织方式完全一致(参见 python/mxnet/gluon/data/vision/datasets.py 的文档字符串与_list_images实现:它会按文件夹枚举类别、生成synsets属性并建立(filename, label)列表)。

使用 Gluon 进行训练

定义超参数

先导入必要依赖:

import math import os import time from mxnet import autograd from mxnet import gluon, init from mxnet.gluon import nn from mxnet.gluon.data.vision import transforms from mxnet.gluon.model_zoo.vision import resnet50_v2

然后定义微调所需的超参数。教程采用 MXNet 学习率调度器在训练过程中动态调整学习率,详细的调度器用法可参考仓库教程 learning_rate_schedules.md。示例中epochs设为 1 仅为快速演示,正式训练请改为 40。

classes = 102 epochs = 1 lr = 0.001 per_device_batch_size = 32 momentum = 0.9 wd = 0.0001 lr_factor = 0.75 # 学习率在这些 epoch 处发生衰减 lr_epochs = [10, 20, 30] num_gpus = mx.context.num_gpus() # 可将 num_workers 替换为设备上的 CPU 核数 num_workers = 8 ctx = [mx.gpu(i) for i in range(num_gpus)] if num_gpus > 0 else [mx.cpu()] batch_size = per_device_batch_size * max(num_gpus, 1)

要点说明:

  • batch_size按 GPU 数量翻倍,保证多卡训练时每个设备仍有per_device_batch_size大小的批次;
  • num_workers控制 DataLoader 的并行读取进程数,建议与 CPU 核数相当;
  • 后续所有训练与验证代码均可同时运行在 CPU(mx.cpu())或 GPU 列表ctx上。

数据增强与 Transform 流水线

训练集较小是微调场景的普遍痛点,数据增强通过对训练图像做轻微改动(模型会视其为不同图像)来扩充有效样本量,有助于提升最终精度。这里结合 Gluon 的 Dataset、DataLoader 与 Transform API,对训练图像依次执行:

  1. 随机裁剪并缩放到 224×224;
  2. 随机水平翻转;
  3. 随机抖动颜色并添加光照扰动;
  4. 将数据从[height, width, num_channels]转置为[num_channels, height, width],并把像素值从[0, 255]映射到[0, 1]
  5. 用 ImageNet 数据集的均值与标准差做归一化。

验证与推理阶段只需执行第 1、4、5 步。同时要把均值/标准差保存为 NDArray 文件,供后续 C++ 推理复用。

jitter_param = 0.4 lighting_param = 0.1 # 归一化图像(值域 0~1)所用的 mean 与 std mean = [0.485, 0.456, 0.406] std = [0.229, 0.224, 0.225] training_transformer = transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomFlipLeftRight(), transforms.RandomColorJitter(brightness=jitter_param, contrast=jitter_param, saturation=jitter_param), transforms.RandomLighting(lighting_param), transforms.ToTensor(), transforms.Normalize(mean, std) ]) validation_transformer = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean, std) ]) # 保存 mean/std 的 NDArray 值,供推理阶段使用 mean_img = mx.nd.stack(*[mx.nd.full((224, 224), m) for m in mean]) std_img = mx.nd.stack(*[mx.nd.full((224, 224), s) for s in std]) mx.nd.save('mean_std_224.nd', {"mean_img": mean_img, "std_img": std_img}) train_path = os.path.join(path, 'train') val_path = os.path.join(path, 'valid') test_path = os.path.join(path, 'test') # 加载数据并应用预处理(transforms) train_data = gluon.data.DataLoader( gluon.data.vision.ImageFolderDataset(train_path).transform_first(training_transformer), batch_size=batch_size, shuffle=True, num_workers=num_workers) val_data = gluon.data.DataLoader( gluon.data.vision.ImageFolderDataset(val_path).transform_first(validation_transformer), batch_size=batch_size, shuffle=False, num_workers=num_workers) test_data = gluon.data.DataLoader( gluon.data.vision.ImageFolderDataset(test_path).transform_first(validation_transformer), batch_size=batch_size, shuffle=False, num_workers=num_workers)

这里transform_first只对数据(图像)做变换而保持标签不变,是图像分类数据加载的推荐写法;RandomColorJitterbrightness/contrast/saturation抖动幅度统一取jitter_param = 0.4RandomLighting扰动系数取0.1

加载预训练模型与 Hybridization

我们使用在 ImageNet(1000 类)上预训练的resnet50_v2。由于花卉数据只有 102 类,必须把最后一层 softmax(输出层)重新定义为 102 维,并重新初始化该层参数。

在训练之前,还要了解 Gluon 的一个独特特性:Hybridization(混合化)。调用net.hybridize()即可把命令式(imperative)代码转换为静态符号图执行,带来两大收益——更优的执行性能、以及更容易序列化以便部署。更深入的解释可参考仓库教程 hybridize.md。

# 从模型动物园加载预训练 resnet50_v2 finetune_net = resnet50_v2(pretrained=True, ctx=ctx) # 类别数不同,替换最后一层 softmax with finetune_net.name_scope(): finetune_net.output = nn.Dense(classes) finetune_net.output.initialize(init.Xavier(), ctx=ctx) # hybridize 以获得更好性能 finetune_net.hybridize() num_batch = len(train_data) # 配置学习率调度器 iterations_per_epoch = math.ceil(num_batch) # 学习率在以下 step 处衰减 lr_steps = [epoch * iterations_per_epoch for epoch in lr_epochs] schedule = mx.lr_scheduler.MultiFactorScheduler(step=lr_steps, factor=lr_factor, base_lr=lr) # 配置带学习率调度器的优化器、评估指标与损失函数 sgd_optimizer = mx.optimizer.SGD(learning_rate=lr, lr_scheduler=schedule, momentum=momentum, wd=wd) metric = mx.metric.Accuracy() softmax_cross_entropy = gluon.loss.SoftmaxCrossEntropyLoss()

从源码结构看,resnet50_v2(python/mxnet/gluon/model_zoo/vision/resnet.py)实际是get_resnet(2, 50)的封装,对应ResNetV2类(resnet.py)。ResNetV2继承自HybridBlock,网络主体由features(含BatchNormConv2D、残差 stage、GlobalAvgPool2DFlatten)与outputnn.Dense(classes))构成,且hybrid_forward只接受符号接口F,这正是它可以被hybridize()编译为静态图、并被export()导出为符号文件的前提。

几个配置要点:

  • MultiFactorScheduler在第 10、20、30 个 epoch 处将学习率乘以factor=0.75
  • lr_steps把「按 epoch 衰减」换算成「按 iteration 衰减」(step = epoch * iterations_per_epoch);
  • SGD使用 momentum=0.9、weight decay=0.0001;损失为SoftmaxCrossEntropyLoss,指标为Accuracy

在自定义数据集上微调

下面定义验证函数并启动微调循环。gluon.utils.split_and_load(实现见 python/mxnet/gluon/utils.py)会把一个批次按ctx列表切分到多张卡上,even_split=False表示最后一个设备可少分数据。

def test(net, val_data, ctx): metric = mx.metric.Accuracy() for i, (data, label) in enumerate(val_data): data = gluon.utils.split_and_load(data, ctx_list=ctx, even_split=False) label = gluon.utils.split_and_load(label, ctx_list=ctx, even_split=False) outputs = [net(x) for x in data] metric.update(label, outputs) return metric.get() trainer = gluon.Trainer(finetune_net.collect_params(), optimizer=sgd_optimizer) # 从 epoch 1 开始,便于学习率计算 for epoch in range(1, epochs + 1): tic = time.time() train_loss = 0 metric.reset() for i, (data, label) in enumerate(train_data): # 取图像与标签 data = gluon.utils.split_and_load(data, ctx_list=ctx, even_split=False) label = gluon.utils.split_and_load(label, ctx_list=ctx, even_split=False) with autograd.record(): outputs = [finetune_net(x) for x in data] loss = [softmax_cross_entropy(yhat, y) for yhat, y in zip(outputs, label)] for l in loss: l.backward() trainer.step(batch_size) train_loss += sum([l.mean().asscalar() for l in loss]) / len(loss) metric.update(label, outputs) _, train_acc = metric.get() train_loss /= num_batch _, val_acc = test(finetune_net, val_data, ctx) print('[Epoch %d] Train-acc: %.3f, loss: %.3f | Val-acc: %.3f | learning-rate: %.3E | time: %.1f' % (epoch, train_acc, train_loss, val_acc, trainer.learning_rate, time.time() - tic)) _, test_acc = test(finetune_net, test_data, ctx) print('[Finished] Test-acc: %.3f' % (test_acc))

训练循环的关键点:

  • autograd.record()记录前向计算图,l.backward()之后调用trainer.step(batch_size)更新参数;
  • 注意trainer.step传入的是全局batch_size(含多卡),而不是单卡批次大小;
  • 每个 epoch 结束后在验证集上评估一次精度,全部完成后在测试集上给出最终指标。

以下为 40 个 epoch 的示例训练输出(教程原文记录):

[Epoch 40] Train-acc: 0.945, loss: 0.354 | Val-acc: 0.955 | learning-rate: 4.219E-04 | time: 17.8 [Finished] Test-acc: 0.952

该结果来自一台配备 4 块 Tesla V100 GPU 的实例:40 个 epoch 约 12 分钟即达到约 95.5% 的测试精度。之所以如此高效,正是因为模型已在约 130 万张图像的 ImageNet 上预训练,对小数据集的特征提取非常有效——这正是迁移学习的核心价值。

保存微调后的模型

训练完成后,用export把模型序列化为模型文件:

finetune_net.export("flower-recognition", epoch=epochs)

export会在当前目录生成两个文件:模型结构文件flower-recognition-symbol.json与参数文件flower-recognition-0040.params0040对应训练的 40 个 epoch,若epochs=1则生成flower-recognition-0001.params),二者即为下一节部署推理的输入。

从源码看,HybridBlock.export(python/mxnet/gluon/block.py)有以下行为值得注意:

  • 必须先调用net.hybridize()并至少前向执行一次,否则会抛出RuntimeError("Please first call block.hybridize() and then run forward with this block at least once before calling export.")——因为导出需要_cached_graph中缓存的符号图;
  • 单输入模型的输入节点名固定为data(多输入时为data0data1…),这正是后续推理时data_shapes=[('data', (1, 3, 224, 224))]的来源;
  • 参数按arg:/aux:前缀分别写入.params文件,供load_checkpointset_params还原。

用 MXNet Module API 加载模型并推理

MXNet 为部署推理提供了多种接口:可以使用 MXNet Model Server 直接托管模型并对外提供服务,也可以借助 Python、Java、Scala、C++ 等多种语言 API 把模型集成进既有服务。本节演示 Python 侧使用 Module API 完成一次预测。

推理整体分为五步:

  1. 加载模型结构(symbol 文件)与训练好的参数(params 文件);
  2. 加载 synset 文件获取类别名称;
  3. 加载图片并应用与训练时验证集相同的变换;
  4. 对图片数据执行一次前向计算;
  5. 把输出概率转换为预测的类别名。
import numpy as np from collections import namedtuple ctx = mx.cpu() # 加载模型 symbol 与 params sym, arg_params, aux_params = mx.model.load_checkpoint('flower-recognition', epochs) mod = mx.mod.Module(symbol=sym, context=ctx, label_names=None) mod.bind(for_training=False, data_shapes=[('data', (1, 3, 224, 224))], label_shapes=mod._label_shapes) mod.set_params(arg_params, aux_params, allow_missing=True) # 加载 synset 以获取类别名 with open('synset.txt', 'r') as f: labels = [l.rstrip() for l in f] # 加载一张待预测图片 img = mx.image.imread('./data/test/lotus/image_01832.jpg') # 应用训练时相同的变换 img = validation_transformer(img) # batchify:扩展为 batch 维度 img = img.expand_dims(axis=0) Batch = namedtuple('Batch', ['data']) mod.forward(Batch([img])) prob = mod.get_outputs()[0].asnumpy() prob = np.squeeze(prob) idx = np.argmax(prob) print('probability=%f, class=%s' % (prob[idx], labels[idx]))

执行结果如下,可见图片被正确分类为 lotus:

probability=9.798435, class=lotus

几个易错点需要留意:

  • 输入形状必须与训练一致data_shapes=[('data', (1, 3, 224, 224))]对应 224×224、RGB 三通道、单样本 batch;若训练时用了其他尺寸,这里需同步修改;
  • 变换必须与验证一致:推理仍使用validation_transformer(Resize 256 → CenterCrop 224 → ToTensor → Normalize),否则归一化域不一致会导致精度骤降;
  • synset 顺序必须与训练一致synset.txt由工具脚本按类别名排序生成,而ImageFolderDataset内部也按排序枚举类别,二者对齐才能正确解析预测序号。

部署路径与后续方向

模型导出为-symbol.json-params后,就可以脱离训练代码独立部署:

  • C++ 部署:继续阅读仓库 cpp-package/example/inference 下的推理示例,了解如何使用 C++ API 加载同一套模型文件、复用mean_std_224.nd完成预处理并执行前向计算;
  • Java / Scala 部署:仓库 scala-package 提供了 Java 与 Scala 的推理示例(Java 示例位于scala-package/examples/src/main/java/org/apache/mxnetexamples/javaapi/infer);
  • 服务化部署:可以使用 MXNet Model Server 启动推理服务,把训练好的模型托管为 HTTP 接口,供上层业务调用。

参考资源

  • docs/tutorial_utils/data/oxford_102_flower_dataset.py:Oxford 102 数据下载与整理脚本(含 102 个花名清单)
  • python/mxnet/gluon/model_zoo/vision/resnet.py:ResNetV2/resnet50_v2实现
  • python/mxnet/gluon/block.py:HybridBlock.export序列化实现
  • learning_rate_schedules.md:学习率调度器详解
  • hybridize.md:Hybridization 原理与用法
  • Gluon 微调相关实践可参考公开的《动手学深度学习》(d2l) 微调章节与 GluonCV 迁移学习教程

【免费下载链接】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

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/21 1:19:41

MATLAB批量处理AFM力曲线:从NSMatlabUtilities到全流程自动化实战

接手过一个接近尾声的材料表征项目,那阵子我每天的工作就是对着 Bruker NanoScope Analysis 里的力曲线,一条一条点开、框选基线、找接触点、拟合、导出。单个文件里 Force Volume 测了 3232 个点,一千多条曲线,再乘以十几个样品&…

作者头像 李华
网站建设 2026/9/21 1:16:46

基于555电路与单片机的DC-AC逆变器设计:C语言实现与调试指南

简介:面向有单片机与电力电子基础的研发人员和技术爱好者,这份基于C语言的直流-交流变换器设计实例,围绕555电路与单片机协同实现逆变输出的项目化学习需求展开。文档完整覆盖硬件电路设计,包括电源管理、555定时器、单片机控制、…

作者头像 李华
网站建设 2026/9/21 1:16:14

Cline vs Roo Code:同一把 TaoToken Key 跑完前端重构任务

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/21 1:10:07

AI工作台核心不神秘:WorkBuddy的产品化、生态与规模工程

前两天一个朋友给我发来一段录屏,说 WorkBuddy 帮他把跨境电商的订单整理流程从一小时压到了十分钟,他反复问我:这里面的核心技术是不是特别深?老实说,这个评价我听得太多了。市面上关于 WorkBuddy 的讨论,…

作者头像 李华
网站建设 2026/9/21 1:06:15

Scissor算法调参实战:alpha与cutoff参数优化指南

1. 为什么Scissor算法的alpha和cutoff值得单独拎出来讲做单细胞数据分析的人,迟早会碰到一个场景:你手里有一份单细胞转录组数据,同时还有一份表型数据(比如生存时间、疾病分组、药物响应),你想知道哪些细胞…

作者头像 李华