MXNet Clojure Module API 实战指南:从 MNIST 训练到模型保存与恢复

发布时间:2026/9/21 19:01:29
MXNet Clojure Module API 实战指南:从 MNIST 训练到模型保存与恢复 MXNet Clojure Module API 实战指南从 MNIST 训练到模型保存与恢复【免费下载链接】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/mxne/mxnet本教程以 MXNet 的 Clojure 绑定org.apache.clojure-mxnet为核心系统讲解 Module API 这一中高层训练接口的完整使用链路如何准备 MNIST 数据迭代器、如何用 Symbol 组装全连接网络、如何通过bind/init-params准备计算、如何用fit/predict/score完成训练、评估与预测以及如何用save-checkpoint/load-checkpoint持久化并恢复模型。读完本文你将能在 Clojure REPL 中独立完成一个神经网络从数据加载、训练到断点续训的完整闭环。本文所讲解的文档位于 docs/static_site/src/pages/api/clojure/docs/tutorials/module.md是 MXNet 官方 Clojure 教程系列index.md中关于 Module API 的核心篇目。仓库中的底层算子实现如FullyConnected、Activation、SoftmaxOutput将为文中的每个 Clojure 调用提供源码级佐证。Module API 是什么Module API 为神经网络计算提供了一套中高层接口它在一个Module对象中同时封装了一个Symbol和一个或多个Executor屏蔽了底层 executor 绑定、内存分配、前向/反向传播调度等细节。因此它同时具备高层接口fit训练、predict预测、score评估适合直接完成端到端流程中间层接口bind、init-params、forward、backward、update等适合需要精细控制训练过程的场景。这种一份 Symbol、多种用途的设计与 MXNet 其他语言绑定如 Python 侧 python/mxnet/model.py 提供的高层 Model/Module 封装一脉相承先构图再绑定数据形状分配显存最后按需训练或推理。准备工作引入依赖命名空间开始之前需要先创建项目的命名空间并引入本教程所需的各个 MXNet Clojure 命名空间(ns docs.module (:require [clojure.java.io :as io] [clojure.java.shell :refer [sh]] [org.apache.clojure-mxnet.eval-metric :as eval-metric] [org.apache.clojure-mxnet.io :as mx-io] [org.apache.clojure-mxnet.module :as m] [org.apache.clojure-mxnet.symbol :as sym] [org.apache.clojure-mxnet.ndarray :as ndarray]))各命名空间的职责如下命名空间别名职责clojure.java.ioio文件存在性检查等 IO 辅助clojure.java.shellsh执行 shell 命令下载数据org.apache.clojure-mxnet.eval-metriceval-metric评估指标如accuracyorg.apache.clojure-mxnet.iomx-io数据迭代器与批处理工具org.apache.clojure-mxnet.modulemModule API 核心org.apache.clojure-mxnet.symbolsym符号计算图构建org.apache.clojure-mxnet.ndarrayndarrayNDArray 张量操作说明Clojure 绑定包clojure-package作为独立包维护本仓库收录的是其官方文档docs/static_site/src/pages/api/clojure/目录以及其底层依赖的 MXNet C 算子实现。文档中所引用的数据下载脚本get_mnist_data.sh位于该 Clojure 包目录下使用时请先完成 Clojure 包的构建与安装。准备数据加载 MNIST 数据集本教程使用经典的 MNIST 手写数字数据集。如果你已经克隆了 MXNet 仓库并进入 Clojure 包目录可以运行辅助脚本自动下载数据(def>(def train-data (mx-io/mnist-iter {:image (str>(let [data (sym/variable data) fc1 (sym/fully-connected fc1 {:data data :num-hidden 128}) act1 (sym/activation relu1 {:data fc1 :act-type relu}) fc2 (sym/fully-connected fc2 {:data act1 :num-hidden 64}) act2 (sym/activation relu2 {:data fc2 :act-type relu}) fc3 (sym/fully-connected fc3 {:data act2 :num-hidden 10}) out (sym/softmax-output softmax {:data fc3})] out) ;#object[org.apache.mxnet.Symbol 0x1f43a406 org.apache.mxnet.Symbol1f43a406]用as-线程宏改写Clojure 开发者也可以使用as-线程宏让数据流方向更加直观——每一层都基于前一层的结果继续构造(def out (as- (sym/variable data) data (sym/fully-connected fc1 {:data data :num-hidden 128}) (sym/activation relu1 {:data data :act-type relu}) (sym/fully-connected fc2 {:data data :num-hidden 64}) (sym/activation relu2 {:data data :act-type relu}) (sym/fully-connected fc3 {:data data :num-hidden 10}) (sym/softmax-output softmax {:data data}))) ; #tutorial.module/out底层算子的参数语义这些 Clojure 调用最终都会路由到 MXNet 底层的 C 算子其参数语义可以直接在算子源码中核实fully-connected的:num-hidden参数对应 src/operator/nn/fully_connected-inl.h 中FullyConnectedParam的num_hidden字段其声明set_lower_bound(1)表明输出隐藏节点数至少为 1同结构体中no_bias默认false是否禁用偏置与flatten默认true是否将除第一维以外的维度展平也都会影响该层的行为。activation的:act-type relu对应 src/operator/nn/activation-inl.h 中ActivationOpType枚举的kReLU该枚举还支持sigmoid、log_sigmoid、mish、tanh、softrelu、softsign等激活类型。softmax-output在 src/operator/softmax_output-inl.h 中注册的算子类型名为SoftmaxOutput它同时完成 softmax 概率计算与交叉熵损失的反向传播其输入依赖声明DeclareBackwardDependency要求同时提供数据输入与标签输入——这正是数据迭代器中:label-name softmax_label需要与之一一对应的原因。准备 Module 进行计算有了 Symbol 之后就可以构造 Module 并进入绑定与参数初始化阶段。默认情况下Module 运行在CPUcontext默认为 CPU如果需要数据并行可以指定 GPU 上下文或一组 GPU 上下文例如(m/module out {:contexts [(context/gpu)]})。(let [mod (m/module out)] (- mod (m/bind {:data-shapes (mx-io/provide-data train-data) :label-shapes (mx-io/provide-label train-data)}) (m/init-params)))这里的两个关键步骤bind根据数据迭代器提供的:data-shapes由mx-io/provide-data从训练集推断和:label-shapes由mx-io/provide-label推断为各个层分配设备内存并构造 Executor。形状信息必须与 Symbol 中各层期望的输入形状匹配。init-params初始化网络参数权重、偏置与辅助状态。如果你已经有现成参数也可以改用set-params直接赋值见下文保存与加载一节。完成这两步后就可以用forward、backward等中间层接口逐步控制计算。但如果你只是想做端到端训练完全不必手动调用bind和init-params——fit函数在需要时会自动替你完成这两步。训练与预测Module 为训练、预测和评估提供了开箱即用的高层 API。训练fit调用fit传入训练数据迭代器即可开始训练同时传入评估数据与轮数fit会在每个 epoch 结束后自动在评估集上打分(def mod (m/fit (m/module out) {:train-data train-data :eval-data test-data :num-epoch 1})) ;; Epoch 0 Train- [accuracy 0.12521666] ;; Epoch 0 Time cost- 8392 ;; Epoch 0 Validation- [accuracy 0.2227]输出日志中依次为当前 epoch 号、训练集上的指标accuracy、该 epoch 耗时、验证集上的指标。fit的行为可以通过:fit-params进一步定制回调函数通过batch-end-callback传入批次结束回调、epoch-end-callback传入轮次结束回调例如每个 epoch 结束时自动保存模型——Python 侧对应的检查点机制可在 python/mxnet/callback.py 的save_checkpoint回调中看到其典型用法save_checkpoint(prefix, iter_no 1, sym, arg, aux)优化器通过optimizer指定优化算法评估指标通过eval-metric指定评估指标。关于fit-params的完整可选参数请查阅 Clojure API 参考中fit-param函数对应的选项说明。预测predict与predict-every-batchpredict接收一个 DataIter返回该数据集上全部样本的预测结果NDArray 的集合(def results (m/predict mod {:eval-data test-data})) (first results) ;#object[org.apache.mxnet.NDArray 0x3540b6d3 org.apache.mxnet.NDArraya48686ec] (first (ndarray/-vec (first results))) ;0.08261358first取出的第一个预测向量通过ndarray/-vec转为 Clojure 向量后其首个元素0.08261358即第一个样本在第一个类别上的预测概率。当预测结果可能大到内存装不下时改用predict-every-batch它逐批次返回预测结果配合mx-io/reduce-batches可以边取边处理(let [preds (m/predict-every-batch mod {:eval-data test-data})] (mx-io/reduce-batches test-data (fn [i batch] (println (str pred is (first (get preds i)))) (println (str label is (mx-io/batch-label batch))) ;;; do something (inc i))))评估score如果只需要在测试集上打分、不需要保留预测输出调用score并传入数据迭代器与评估指标(m/score mod {:eval-data test-data :eval-metric (eval-metric/accuracy)}) ;[accuracy 0.2227]score会遍历数据迭代器中的每个批次执行预测并用给定的评估指标计算得分返回值为[指标名 得分]形式的向量例如[accuracy 0.2227]。评估结果同时保存在eval-metric对象内部之后可以随时查询。保存与加载模型保存检查点save-checkpoint在每个训练 epoch 结束时可以用save-checkpoint保存模块参数必要时连同优化器状态一并保存(let [save-prefix my-model] (doseq [epoch-num (range 3)] (mx-io/do-batches train-data (fn [batch ;; do something ])) (m/save-checkpoint mod {:prefix save-prefix :epoch epoch-num :save-opt-states true}))) ;; INFO org.apache.mxnet.module.Module: Saved checkpoint to my-model-0000.params ;; INFO org.apache.mxnet.module.Module: Saved optimizer state to my-model-0000.states ;; INFO org.apache.mxnet.module.Module: Saved checkpoint to my-model-0001.params ;; INFO org.apache.mxnet.module.Module: Saved optimizer state to my-model-0001.states ;; INFO org.apache.mxnet.module.Module: Saved checkpoint to my-model-0002.params ;; INFO org.apache.mxnet.module.Module: Saved optimizer state to my-model-0002.states这里的关键参数是:prefix模型文件名的前缀。my-model会生成my-model-0000.params这样的文件其中 4 位数字为 epoch 编号:epoch当前 epoch 编号决定文件名中的编号段:save-opt-states true同时保存优化器状态momentum 等这是恢复训练所必需的。保存结果会生成两类文件*.params网络参数与*.states优化器状态。加载检查点load-checkpointload-checkpoint从磁盘恢复一个已保存的 Module:epoch指定要加载哪个 epoch 的检查点(def new-mod (m/load-checkpoint {:prefix my-model :epoch 1 :load-optimizer-states true})) new-mod ; #object[org.apache.mxnet.module.Module 0x5304d0f4 org.apache.mxnet.module.Module5304d0f4]加载得到的 Module 尚未绑定内存、也未初始化参数因此需要像第一次使用时那样执行bind与init-params(- new-mod (m/bind {:data-shapes (mx-io/provide-data train-data) :label-shapes (mx-io/provide-label train-data)}) (m/init-params))读取与写入参数params、set-paramsparams以[arg-params aux-params]二元组的形式返回当前全部参数arg-params是参数各层权重与偏置的名称 → NDArray映射aux-params是辅助状态如 BatchNorm 的均值和方差的映射(let [[arg-params aux-params] (m/params new-mod)] {:arg-params arg-params :aux-params aux-params}) ;; {:arg-params ;; {fc3_bias ;; #object[org.apache.mxnet.NDArray 0x39adc3b0 org.apache.mxnet.NDArray49caf426], ;; fc2_weight ;; #object[org.apache.mxnet.NDArray 0x25baf623 org.apache.mxnet.NDArraya6c8f9ac], ;; fc1_bias ;; #object[org.apache.mxnet.NDArray 0x6e089973 org.apache.mxnet.NDArray9f91d6eb], ;; fc3_weight ;; #object[org.apache.mxnet.NDArray 0x756fd109 org.apache.mxnet.NDArray2dd0fe3c], ;; fc2_bias ;; #object[org.apache.mxnet.NDArray 0x1dc69c8b org.apache.mxnet.NDArrayd128f73d], ;; fc1_weight ;; #object[org.apache.mxnet.NDArray 0x20abc769 org.apache.mxnet.NDArrayb8e1c5e8]}, ;; :aux-params {}}可以看到本示例网络的参数名遵循fcN_weight/fcN_bias的命名约定与本教程 Symbol 构建时的层名fc1、fc2、fc3一一对应没有 BatchNorm 层因此:aux-params为空映射。要向模块写入参数与辅助状态使用set-params例如把另一个模块的参数拷贝过来(m/set-params new-mod {:arg-params (m/arg-params new-mod) :aux-params (m/aux-params new-mod)}) ; #object[org.apache.mxnet.module.Module 0x5304d0f4 org.apache.mxnet.module.Module5304d0f4]从检查点恢复训练断点续训要基于保存的检查点继续训练核心是让fit不再随机初始化参数并告诉它从哪个 epoch 继续用load-checkpoint加载已保存参数上文new-mod由于数据迭代器已被消费过调用fit前必须先reset训练集与测试集否则会报错构造fit-params通过begin-epoch指定起始 epoch即上次保存的 epoch 编号fit就知道从该 epoch 接着训练且不会重新随机初始化参数。;; reset the training data before calling fit or you will get an error (mx-io/reset train-data) (mx-io/reset test-data) (m/fit new-mod {:train-data train-data :eval-data test-data :num-epoch 2 :fit-params (- (m/fit-params {:begin-epoch 1}))})下一步学习阅读 Symbolic API 教程学习如何用层组装神经网络计算图阅读 NDArray API 教程掌握向量/矩阵/张量级操作阅读 KVStore API 教程了解如何借助 KVStore 进行多 GPU 与多主机分布式训练完整的 Clojure 教程索引见 Clojure 教程主页。【免费下载链接】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/mxne/mxnet创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考