ARTICLE DETAIL

资讯详情

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

Apache MXNet Scala Module API 实战指南:从 Symbol 到训练、预测与模型持久化的完整流程

Apache MXNet Scala Module API 实战指南:从 Symbol 到训练、预测与模型持久化的完整流程 Apache MXNet Scala Module API 实战指南从 Symbol 到训练、预测与模型持久化的完整流程【免费下载链接】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 Scala 包scala-package的开发者系统讲解Module API这一套位于底层Executor之上的中高级编程接口。你将从零构造一个 MLP 符号网络学会如何通过bind()与initParams()把网络激活为可计算模块再借助fit()、predict()、score()完成端到端的训练、推理与评估最后掌握saveCheckpoint/loadCheckpoint的断点续训方案让训练中途断电不再前功尽弃。阅读本指南后你将能够独立编写可运行的 Scala 训练脚本并理解每个调用背后的源码机制。一、认识 Module APIModule、Symbol 与 Executor 的关系在 MXNet Scala 中Module API 为神经网络计算提供了一层介于底层与高层之间的接口。核心概念是一个module是BaseModule子类的一个实例最常用的类是Module。Module包装了一个Symbol和一个或多个Executor。这三者的分工可以这样理解Symbol定义网络的计算图是什么结构例如一个三层全连接 MLPExecutor将 Symbol 绑定到具体数据形状并分配显存/内存后得到的可执行实例怎么算Module把两者包装起来对外暴露训练、预测、评估、保存/加载等完整生命周期操作怎么用。从源码看BaseModule.scala 的类注释明确描述了一个模块应具备的多阶段状态机Initial state初始态尚未分配内存不可计算Binded已绑定输入、输出、参数形状全部已知内存已分配可开始计算Parameter initialized参数已初始化未初始化参数就进行计算会产生未定义输出Optimizer installed优化器已安装安装优化器后前向-反向得到的梯度才能驱动参数更新。BaseModule还定义了模块之间交互所需的协议信息dataNames、outputNames绑定前即可报告以及绑定后的dataShapes、labelShapes、outputShapes和getParams/setParams等。理解了这套状态机后面所有 API 的调用顺序先bind再initParams就顺理成章了。所有 Module API 都位于org.apache.mxnet.module包下。BaseModule的子类除Module外还包括支持变长序列的BucketingModule、可将多个模块链式组合的SequentialModule本指南以最常用的Module为主线展开。二、准备一个可用于计算的 Module2.1 构造 Module从一个 Symbol 开始Module类的构造函数接受一个Symbol作为输入。下面用一个三层 MLP 作为示例构建过程与 MXNet 其他语言前端一致import org.apache.mxnet._ import org.apache.mxnet.module.{FitParams, Module} // 构造一个简单的 MLP val data Symbol.Variable(data) val fc1 Symbol.api.FullyConnected(Some(data), num_hidden 128, name fc1) val act1 Symbol.api.Activation(Some(fc1), relu, relu1) val fc2 Symbol.api.FullyConnected(Some(act1), num_hidden 64, name fc2) val act2 Symbol.api.Activation(Some(fc2), relu, relu2) val fc3 Symbol.api.FullyConnected(Some(act2), num_hidden 10, name fc3) val out Symbol.api.SoftmaxOutput(fc3, name softmax) // 构造 module val mod new Module(out)在仓库示例 MnistMlp.scala 中同样的网络使用了等价的符号式函数调用风格Symbol.FullyConnected(...)(...)(Map(...))二者构造出的计算图完全一致可按个人习惯选择。2.2 构造函数的关键参数从 Module.scala 的类定义可以看到Module构造函数除symbolVar外还支持以下参数参数默认值说明dataNamesIndexedSeq(data)输入数据data的名称列表labelNamesIndexedSeq(softmax_label)标签label的名称列表若网络不需要标签可传null或空序列contextsArray(Context.cpu())计算设备上下文默认 CPUworkLoadListNone各设备上的工作量分配比例默认均匀分配长度必须与contexts一致fixedParamNamesNone需要固定不参与训练的参数名集合常用于迁移学习冻结底层默认情况下context是 CPU。如果需要数据并行可以传入一个 GPU context 或 GPU context 数组Context.gpu(0)、Context.gpu(1)等。Module同时提供了配套的 Builder 模式支持setContext、setDataNames、setLabelNames、setWorkLoadList、setFixedParamNames链式调用后再build()适合构造参数较多的场景。2.3 bind 与 initParams让模块通电刚构造出来的 Module 还处于初始态——没有分配任何内存。开始计算前必须依次执行两步bind()根据数据形状分配设备内存构建 ExecutorinitParams()初始化参数权重与辅助状态auxiliary states。bind()的数据形状通常直接取自DataIter的provideData/provideLabelmod.bind(dataShapes train_dataiter.provideData, labelShapes Some(train_dataiter.provideLabel)) mod.initParams()提示如果只是想简单地拟合一个模块可以跳过显式的bind()和initParams()因为fit()内部会在需要时自动调用它们见 BaseModule.scala 中fit的实现。完成这两步后模块即进入参数已初始化状态可以用forward()、backward()等函数进行计算了。bind()还有一些值得了解的参数Module.scalaforTraining默认true决定 Executor 是否以训练模式绑定inputsNeedGrad默认false是否需要计算对输入数据的梯度实现模块组合时可能需要forceRebind默认false为true时强制重新绑定常用于从训练切换到推理sharedModule用于 bucketing 场景共享一组参数gradReq默认write梯度累积方式可选write、add、null。三、训练、预测与评估3.1 高层训练接口fit()fit()是模块提供的最顶层训练 API输入一个或多个DataIter即可完成整个训练流程import org.apache.mxnet.optimizer.SGD val mod new Module(softmax) mod.fit(train_dataiter, evalData scala.Option(eval_dataiter), numEpoch n_epoch, fitParams new FitParams() .setOptimizer(new SGD(learningRate 0.1f, momentum 0.9f, wd 0.0001f)))从源码看fit()的执行流程BaseModule.scala依次为按训练数据形状自动bind(...)forTraining true按fitParams中的initializer自动initParams(...)按kvstore与optimizer自动initOptimizer(...)进入 epoch 循环对每个 batch 执行forwardBackward(dataBatch)→update()→updateMetric(...)epoch 结束时在验证集上调用score(...)并打印训练/验证指标。这个接口与旧版FeedForward类的用法非常相似方便老用户平滑迁移。fit()还支持通过setBatchEndCallback传入 batch 级回调、setEpochEndCallback传入 epoch 级回调用setOptimizer、setEvalMetric等设置训练细节。3.2 FitParams训练配置的集中地FitParams是fit()的配置载体BaseModule.scala其常用 setter 及默认值如下Setter 方法默认值作用setEvalMetricnew Accuracy()训练过程中显示的评估指标setValidationMetricNone验证集专用指标缺省时复用evalMetricsetOptimizernew SGD()参数更新优化器setKVStorelocalKVStore 类型local/dist_sync/dist_async等setInitializernew Uniform(0.01f)参数初始化器setArgParams/setAuxParamsnull已有参数/辅助状态断点续训时使用setAllowMissingfalse是否允许参数缺失并用初始化器补齐setForceRebindfalse是否强制重新绑定 ExecutorsetForceInitfalse是否强制重新初始化参数setBeginEpoch0起始 epoch 编号续训时为上次保存的 epoch 1setBatchEndCallback/setEpochEndCallbackNonebatch/epoch 结束回调setEvalEndCallback/setEvalBatchEndCallbackNone评估阶段回调setMonitorNone计算监控器这些 setter 全部返回FitParams自身支持链式调用。更多细节可查阅org.apache.mxnet.module.FitParams的 API 文档。3.3 用 predict() 做预测predict()接收一个DataIter模块会遍历其中全部 batch 并收集、返回所有预测结果mod.predict(val_dataiter)从实现看predict(evalData)内部先调用predictEveryBatch逐批预测再将各 batch 的输出按输出序号拼接concatenate成IndexedSeq[NDArray]BaseModule.scala。返回值格式的详细说明可参考org.apache.mxnet.module.BaseModule的 API 文档。注意predict(DataIter)会尝试把各 batch 的输出合并因此要求每个 batch 的输出数量一致若网络输出数量随 batch 变化如 bucketing合并会失败——此时应改用下面的predictEveryBatch。3.4 内存受限时使用 predictEveryBatch当预测结果可能大到无法全部装入内存时请使用predictEveryBatchAPI。它逐 batch 返回预测结果嵌套结构IndexedSeq[IndexedSeq[NDArray]]配合数据迭代器逐个 batch 处理val preds mod.predictEveryBatch(val_dataiter) val_dataiter.reset() var i 0 while (val_dataiter.hasNext) { val batch val_dataiter.next() val predLabel: Array[Int] NDArray.argmax_channel(preds(i)(0)).toArray.map(_.toInt) val label batch.label(0).toArray.map(_.toInt) // do something... i 1 }predictEveryBatch的返回结构形如[ [out1_batch1, out2_batch1, ...], [out1_batch2, out2_batch2, ...] ]即外层为 batch、内层为该 batch 的各个输出。仓库示例 MnistMlp.scala 展示了用该接口逐批计算验证集准确率的完整写法。3.5 用 score() 只评估不出预测如果只需要在测试集上评估、不需要预测输出调用score()并传入一个DataIter和一个EvalMetricmod.score(val_dataiter, metric)score()会对DataIter中的每个 batch 执行前向计算并用给定的EvalMetric累计评估分数评估结果保存在metric对象中事后可查询。其源码实现BaseModule.scala还支持numBatch限制评估的 batch 数、reset评估前是否重置迭代器、batchEndCallback/scoreEndCallback评估回调等参数。在 MnistMlp.scala 中可以看到mod.score(test, new Accuracy).get取回(名称, 数值)的用法。四、保存与加载模块参数4.1 训练过程中保存 checkpoint使用 checkpoint 回调可以在每个训练 epoch 保存模块参数。也可以像下面的代码一样在自定义训练循环中手动保存val modelPrefix: String mymodel for (epoch - 0 until 5) { while (train_dataiter.hasNext) { // forward backward pass // do something... } val checkpoint mod.saveCheckpoint(modelPrefix, epoch, saveOptStates true) }从源码看saveCheckpointModule.scala会同时产出三份文件$prefix-symbol.json网络结构Symbol 图$prefix-%04d.params参数文件如mymodel-0003.params$prefix-%04d.states优化器状态文件仅当saveOptStates true用于无缝续训。参数文件的内部组织在 BaseModule.scala 的saveParams中有体现arg 参数以arg:名称为键、辅助状态以aux:名称为键统一写入NDArray.save加载时loadParams则按arg:/aux:前缀反解析并调用setParams回填BaseModule.scala。4.2 从 checkpoint 加载模块加载已保存的模块参数使用loadCheckpoint工厂方法val mod Module.loadCheckpoint(modelPrefix, loadModelEpoch, loadOptimizerStates true)Module.loadCheckpointModule.scala内部调用Model.loadCheckpoint(prefix, epoch)读取符号与参数构造出新的Module实例并直接标记paramsInitialized true若指定loadOptimizerStates true还会预载$prefix-%04d.states中的优化器状态使续训时的动量等状态得以保留。4.3 初始化、获取与设置参数初始化参数先bind构造 Executor再调用initParams()mod.bind(dataShapes train_dataiter.provideData, labelShapes Some(train_dataiter.provideLabel)) mod.initParams()获取当前参数使用getParams返回(argParams, auxParams)两个名称 → NDArray映射val (argParams, auxParams) mod.getParams注意getParams返回的是 CPU 上的副本参数真正的计算参数可能位于 GPU 等设备上。当paramsDirty标志为真时getParams会先从设备同步最新参数Module.scala。设置参数使用setParams赋值参数与辅助状态mod.setParams(argParams, auxParams)setParams底层委托给initParamsBaseModule.scala支持allowMissing、forceInit、allowExtra等精细控制。4.4 从 checkpoint 恢复训练从保存的 checkpoint 恢复训练时不要调用setParams()而是直接把加载的参数传给fit()让fit()从这些参数出发而不是随机初始化val (argParams, auxParams) mod.getParams // 或从 loadCheckpoint 获得 mod.fit(..., fitParams new FitParams() .setArgParams(argParams) .setAuxParams(auxParams) .setBeginEpoch(beginEpoch))这里的关键是创建FitParams对象后调用setBeginEpoch()传入beginEpoch即上次训练结束的 epoch 编号fit()就能从该 epoch 继续而不是从头开始。从 BaseModule.scala 的注释看beginEpoch的惯例是若此前训练保存于 epoch N则续训时该值应设为 N1。仓库测试 ModuleSuite.scala 完整演示了saveCheckpoint保存 →loadCheckpoint加载含优化器状态的往返流程可作为实战参考。五、进阶其他 BaseModule 子类除了Moduleorg.apache.mxnet.module包还提供两个实用子类详见 module 目录BucketingModule面向变长输入如不同长度的 RNN 序列——同一组参数对应多个不同 Symbolbucket通过switchBucket在它们之间切换forward时自动按 batch 的 bucket 键选择对应的计算图BucketingModule.scala。SequentialModule一个容器模块可通过add(mod1).add(mod2, (take_labels, true), (auto_wiring, true))把多个模块串联成链SequentialModule.scala。示例 SequentialModuleEx.scala 展示了将不含损失的前半网络与含 Softmax 损失的后半网络拼接的写法。其类注释也提醒这类命令式容器在灵活性与效率上不如纯符号图适合作为便捷工具使用。六、下一步学习路线掌握了 Module API 之后可以继续深入 MXNet Scala 的其他核心接口Model API另一种更简单的训练高层接口旧FeedForward的替代Symbolic API用符号算子组装神经网络的计算图IO Data Loading API数据的解析与加载NDArray API向量/矩阵/张量运算KVStore API多 GPU 与多机分布式训练。实际动手时可以直接运行仓库中的 MnistMlp.scala 示例——它同时演示了中间层 API手动bind/initParams/initOptimizer 循环forward/backward/update与高层 APIfit/predict/predictEveryBatch/score两条路线是理解 Module 生命周期的最佳入门代码。【免费下载链接】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),仅供参考
返回列表