MXNet Clojure API 实战指南:张量计算、符号式编程与分布式训练
2026/9/21 15:46:08 网站建设 项目流程
  • 人工智能
  • 深度学习
  • 机器学习

【免费下载链接】mxnet

Lightweight, 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 语言绑定,让 Clojure 开发者能够直接在 JVM 生态中调用 MXNet 的高性能张量计算内核与 GPU 加速能力,并支持构造、定制和训练深度学习模型。本文以仓库中的 Clojure API 文档 为主线,结合仓库内的 NDArray 教程、Symbolic API 教程、Module API 教程 与 KVStore 教程,完整讲解从张量运算、计算图构建到模型训练与多机多卡分布式参数同步的 Clojure 实战方案。读完本文,你将掌握在 Clojure 中完成端到端深度学习任务的完整技术栈。

MXNet 的 Clojure 绑定:能力与定位

MXNet 官方支持 Clojure 编程语言。Clojure API 文档 明确指出,MXNet Clojure 包将灵活高效的 GPU 计算与前沿深度学习能力带入 Clojure 世界:它允许你在纯 Clojure 中编写无缝的、支持多 GPU 的张量/矩阵计算,同时能够构造和定制当前主流的深度学习模型,并将其应用到图像分类、数据科学竞赛等实际任务中。

从架构上看,Clojure 绑定与 Python 绑定 共享同一套 C++ 核心(src/目录下的算子实现,如 src/operator 中丰富的神经网络层算子),通过 JNI 桥接org.apache.mxnet命名空间的 Java 对象(如NDArraySymbolContext),再经由 Clojure 封装层org.apache.clojure-mxnet.*提供函数式 API。因此,你在 Clojure 中创建的模型与张量,与 MXNet 其他语言绑定共用同一套序列化格式和底层执行引擎,模型可以在不同语言绑定之间无缝迁移。

张量与矩阵计算:从零开始

创建 NDArray

Clojure 绑定中的核心张量类型是org.apache.mxnet.NDArray,其操作风格与numpy.ndarray非常相似。首先引入命名空间:

(ns docs.ndarray (:require [org.apache.clojure-mxnet.ndarray :as ndarray] [org.apache.clojure-mxnet.context :as context]))

创建 NDArray 的方式与 NumPy 一致,支持全零、全一以及由向量指定形状构造:

(def a (ndarray/zeros [100 50])) ;; 100 x 50 的全零数组 (def b (ndarray/ones [256 32 128 1])) ;; 四维全一数组 (def c (ndarray/array [1 2 3 4 5 6] [2 3])) ;; 内容为 1..6、形状 2 x 3

其中ndarray/array的第一个参数是扁平数据,第二个参数是目标形状,与 NumPy 的reshape语义对应。

NDArray 还提供便捷的转换与查询接口:

(ndarray/->vec c) ;=> [1.0 2.0 3.0 4.0 5.0 6.0] (ndarray/shape c) ;=> #object[org.apache.mxnet.Shape 0x583c865 "(2,3)"] (ndarray/shape-vec c) ;=> [2 3]

shape返回Shape对象,shape-vec返回纯 Clojure 向量,便于后续编程。

算术运算与就地操作

NDArray 重载了算术运算符,返回新张量且不修改原张量(纯函数语义):

(def a (ndarray/ones [1 5])) (def b (ndarray/ones [1 5])) (-> (ndarray/+ a b) (ndarray/->vec)) ;=> [2.0 2.0 2.0 2.0 2.0] ;; 原数组保持不变 (ndarray/->vec a) ;=> [1.0 1.0 1.0 1.0 1.0] ;; 就地运算符(修改原张量) (ndarray/+= a b) (ndarray/->vec a) ;=> [2.0 2.0 2.0 2.0 2.0]

其他算术运算(ndarray/-ndarray/*ndarray//等)语义完全类似,既有无副作用的版本,也有+=*=这类就地版本,方便在不同场景下权衡内存与不可变性。

切片操作

切片基于轴进行,支持单参数(从第 n 行开始)与双参数(行区间)两种形式:

(def a (ndarray/array [1 2 3 4 5 6] [3 2])) (def a1 (ndarray/slice a 1)) (ndarray/shape-vec a1) ;=> [1 2] (ndarray/->vec a1) ;=> [3.0 4.0] (def a2 (ndarray/slice a 1 3)) (ndarray/shape-vec a2) ;=> [2 2] (ndarray/->vec a2) ;=> [3.0 4.0 5.0 6.0]

矩阵点乘

(def arr1 (ndarray/array [1 2] [1 2])) (def arr2 (ndarray/array [3 4] [2 1])) (def res (ndarray/dot arr1 arr2)) (ndarray/shape-vec res) ;=> [1 1] (ndarray/->vec res) ;=> [11.0]

保存与加载 NDArray

ndarray/save支持将 NDArray 的列表或字典保存到本地文件系统,并且原生支持s3://hdfs://路径,跨语言绑定共享同一格式:

(ndarray/save "filename" {"arr1" arr1 "arr2" arr2}) ;; 也可以使用 "s3://path" 或 "hdfs://path"

加载时返回键值映射:

(def from-file (ndarray/load "filename")) from-file ;=> {"arr1" #object["org.apache.mxnet.NDArray@43d85753"], ; "arr2" #object["org.apache.mxnet.NDArray@5c93def4"]}

多设备支持

设备信息存放在mxnet.Context结构中。创建 NDArray 时可通过:ctx参数指定设备(默认 CPU):

(def cpu-a (ndarray/zeros [100 200])) (ndarray/context cpu-a) ;=> #object[org.apache.mxnet.Context 0x3f376123 "cpu(0)"] (def gpu-b (ndarray/zeros [100 200] {:ctx (context/gpu 0)})) ;; GPU 上创建

context/gpu 0表示第 0 号 GPU。跨设备运算时 MXNet 引擎会自动处理数据搬迁,这也是后续多 GPU 数据并行训练的基础。

Symbolic API:构建计算图

符号组合

Symbolic API 提供了配置计算图的方式,既可以在神经网络层级别组合,也可以做细粒度算子组合。下面是一个经典的两层全连接网络:

(ns docs.symbol (:require [org.apache.clojure-mxnet.executor :as executor] [org.apache.clojure-mxnet.ndarray :as ndarray] [org.apache.clojure-mxnet.symbol :as sym] [org.apache.clojure-mxnet.context :as context])) (def data (sym/variable "data")) (def fc1 (sym/fully-connected "fc1" {:data data :num-hidden 128})) (def act1 (sym/activation "act1" {:data fc1 :act-type "relu"})) (def fc2 (sym/fully-connected "fc2" {:data act1 :num-hidden 64})) (def net (sym/softmax-output "out" {:data fc2}))

sym/variable创建输入占位节点,fully-connected指定隐藏单元数:num-hiddenactivation通过:act-type指定激活类型。利用 Clojure 的as->线程宏,可以将同样的构建过程写成更紧凑的流水线形式:

(as-> (sym/variable "data") data (sym/fully-connected "fc1" {:data data :num-hidden 128}) (sym/activation "act1" {:data data :act-type "relu"}) (sym/fully-connected "fc2" {:data data :num-hidden 64}) (sym/softmax-output "out" {:data data}))

符号同样重载了基本算术运算符,下面创建一个对两个输入求和的计算图:

(def a (sym/variable "a")) (def b (sym/variable "b")) (def c (sym/+ a b))

更复杂的组合与多输入

fully-connected的输入可以是任意符号表达式,例如先做逐元素相加再送入全连接层。通过sym/list-arguments可以查看计算图的所有自由变量(含自动生成的权重与偏置):

(def lhs (sym/variable "data1")) (def rhs (sym/variable "data2")) (def net (sym/fully-connected "fc1" {:data (sym/+ lhs rhs) :num-hidden 128})) (sym/list-arguments net) ;=> ["data1" "data2" "fc1_weight" "fc1_bias"]

分组多个输出

多损失层的网络可以用sym/group将多个输出符号打包成一个计算图,例如同时输出 softmax 分类损失与线性回归损失:

(def net (sym/variable "data")) (def fc1 (sym/fully-connected {:data net :num-hidden 128})) (def net2 (sym/activation {:data fc1 :act-type "relu"})) (def out1 (sym/softmax-output {:data net2})) (def out2 (sym/linear-regression-output {:data net2})) (def group (sym/group [out1 out2])) (sym/list-outputs group) ;=> ["softmaxoutput0_output" "linearregressionoutput0_output"]

sym/list-outputs返回所有输出节点的名称,便于后续按名取结果。

序列化:保存与加载

符号以 JSON 格式保存,格式跨语言、跨云平台通用(支持本地文件与 S3)。通过sym/to-json可直接拿到 JSON 字符串,用于比较两个符号是否等价:

(def a (sym/variable "a")) (def b (sym/variable "b")) (def c (sym/+ a b)) (sym/save c "symbol-c.json") (def c2 (sym/load "symbol-c.json")) (= (sym/to-json c) (sym/to-json c2)) ;=> true

执行符号:bind 与 forward

执行符号需要先用sym/bind将自由变量映射到具体的 NDArray 上,得到Executorexecutor/forward执行前向计算,executor/outputs取回所有输出:

(def ex (sym/bind c {"a" (ndarray/ones [2 2]) "b" (ndarray/ones [2 2])})) (-> (executor/forward ex) (executor/outputs) (first) (ndarray/->vec)) ;=> [2.0 2.0 2.0 2.0]

bind也接受设备上下文参数,实现在 GPU 上运行(前提是已引入对应的 native library jar 依赖;仅使用 CPU 时把gpu_device换成cpu即可):

(def ex (sym/bind c (context/gpu 0) {"a" (ndarray/ones [2 2]) "b" (ndarray/ones [2 2])}))

图解:从组合到执行

关于符号构建、bind、forward、多输出绑定、梯度计算与辅助状态(auxiliary state)的完整流程,仓库提供了带图解说明的 Symbolic Configuration and Execution in Pictures 教程。其核心要点包括:

  • 符号(Symbol)是对计算的描述,构建 API 生成计算图;
  • bind将 NDArray 绑定到参数节点以得到ExecutorExecutor.forward产出结果;
  • mx.symbol.Group分组后绑定可同时获得多个输出,但只绑定你需要的部分,以便系统做更多优化;
  • bind中可指定存放梯度的 NDArray,forward之后调用backward即可得到对应梯度;
  • simple_bind只需给出输入数据形状,自动完成参数分配与 Executor 绑定;
  • 辅助状态(auxiliary states)与参数类似,但不参与梯度计算,常用于跟踪非可导的运行信息。

Module API:训练与推理的高层封装

准备数据

Module API 提供了进行神经网络计算的中高层接口,内部封装一个 Symbol 与一个或多个 Executor。训练以 MNIST 为例:在仓库根目录下执行scripts/get_mnist_data.sh(即教程中cd contrib/clojure-package后运行的脚本)下载数据,再用mx-io/mnist-iter构建数据迭代器:

(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])) (def>(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})))

默认context为 CPU;需要数据并行时,可通过(m/module out {:contexts [(context/gpu)]})指定单个或一组 GPU 上下文。

计算之前需要先bind分配设备内存,并用init-params(或set-params)初始化参数;如果直接使用fit,这些步骤会被自动调用:

(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)))

训练、预测与评估

调用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]

fit通过fit-params支持丰富配置::batch-end-callback/:epoch-end-callback传入批结束/轮结束回调,:optimizer设置优化器,:eval-metric设置评估指标等。

预测用predict,返回所有预测结果的 NDArray 集合:

(def results (m/predict mod {:eval-data test-data})) (first (ndarray/->vec (first results))) ;=>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))) (inc i))))

如果只需要评估指标而不需要预测输出,用score

(m/score mod {:eval-data test-data :eval-metric (eval-metric/accuracy)}) ;=>["accuracy" 0.2227]

评估结果会保存在传入的eval-metric对象中,便于后续查询。

检查点保存与加载

训练过程中用save-checkpoint按 epoch 保存模型参数与优化器状态:

(let [save-prefix "my-model"] (doseq [epoch-num (range 3)] (mx-io/do-batches train-data (fn [batch])) (m/save-checkpoint mod {:prefix save-prefix :epoch epoch-num :save-opt-states true}))) ;; INFO ...: Saved checkpoint to my-model-0000.params ;; INFO ...: Saved optimizer state to my-model-0000.states ;; ... 依次保存到 my-model-0002

加载检查点用load-checkpoint,随后bind+init-params恢复可计算状态:

(def new-mod (m/load-checkpoint {:prefix "my-model" :epoch 1 :load-optimizer-states true})) (-> new-mod (m/bind {:data-shapes (mx-io/provide-data train-data) :label-shapes (mx-io/provide-label train-data)}) (m/init-params))

查看当前参数用params(返回[arg-params aux-params]),例如恢复出的模型包含fc1_weightfc1_biasfc2_weightfc2_biasfc3_weightfc3_bias六组可训练参数:

(let [[arg-params aux-params] (m/params new-mod)] {:arg-params arg-params :aux-params aux-params})

手动赋值参数与辅助状态用set-params

(m/set-params new-mod {:arg-params (m/arg-params new-mod) :aux-params (m/aux-params new-mod)})

从检查点恢复训练

要恢复训练,先重置数据迭代器,再通过fit-params设置begin-epochfit会跳过随机初始化、从保存的 epoch 继续:

(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}))})

KVStore API:多设备与分布式训练

KVStore 提供跨设备(GPU/CPU)与跨主机的键值参数同步能力,是 MXNet 分布式训练的核心组件。

基本 Push 与 Pull

创建本地 KVStore,初始化(key, NDArray)对并拉取:

(ns docs.kvstore (:require [org.apache.clojure-mxnet.kvstore :as kvstore] [org.apache.clojure-mxnet.ndarray :as ndarray] [org.apache.clojure-mxnet.context :as context])) (def kv (kvstore/create "local")) ;; 创建本地 kvstore (def shape [2 3]) ;; 用 key 向量与 ndarray 向量初始化 (kvstore/init kv ["3"] [(ndarray/* (ndarray/ones shape) 2)]) (def a (ndarray/zeros shape)) (kvstore/pull kv ["3"] [a]) (ndarray/->vec a) ;=> [2.0 2.0 2.0 2.0 2.0 2.0]

Push 的聚合语义

对已初始化的 key 可以 push 同形状的新值。push 支持将多个设备上的值推入同一 key,KVStore 会先求和再推送聚合值。下面的例子用 3 个 CPU 各推一个全一数组,最终聚合结果全为 3:

(def cpus [(context/cpu 0) (context/cpu 1) (context/cpu 2)]) (def b [(ndarray/ones shape {:ctx (nth cpus 0)}) (ndarray/ones shape {:ctx (nth cpus 1)}) (ndarray/ones shape {:ctx (nth cpus 2)})]) (kvstore/push kv ["3" "3" "3"] b) (kvstore/pull kv "3" a) (ndarray/->vec a) ;=> [3.0 3.0 3.0 3.0 3.0 3.0]

push时数据可以存放在任意设备上;数据并行训练中,各 worker 的梯度正是通过这种「先聚合后推送」的方式同步到全局参数。

一次 Pull 到多设备

与 push 对称,pull 也可一次将值拉取到多个设备:

(def b [(ndarray/ones shape {:ctx (context/cpu 0)}) (ndarray/ones shape {:ctx (context/cpu 1)})]) (kvstore/pull kv ["3" "3"] b) (map ndarray/->vec b) ;=> ([3.0 ...] [3.0 ...])

批量键值对操作

KVStore 支持对一组 key 同时 init/push/pull,便于管理大规模参数集合:

(def ks ["5" "7" "9"]) (kvstore/init kv ks [(ndarray/ones shape) (ndarray/ones shape) (ndarray/ones shape)]) (kvstore/push kv ks [(ndarray/ones shape) (ndarray/ones shape) (ndarray/ones shape)]) (def b [(ndarray/zeros shape) (ndarray/zeros shape) (ndarray/zeros shape)]) (kvstore/pull kv ks b) (map ndarray/->vec b) ;=> ([1.0 ...] [1.0 ...] [1.0 ...])

学习路径与仓库资源

围绕 Clojure API,仓库提供了一套完整的渐进式学习资料:

  • Clojure API 主页:张量计算的快速上手;
  • NDArray 教程:张量创建、运算、切片、点乘、存取与多设备;
  • Symbolic API 教程:计算图构建、分组、序列化与执行;
  • Module API 教程:从 MNIST 数据准备到训练、预测、检查点管理的完整流程;
  • KVStore 教程:多 GPU/多主机分布式参数同步。

推荐的进阶路线是:先用 NDArray 教程 熟悉张量基础,再通过 Symbolic API 教程 理解计算图(可搭配 图解教程),随后用 Module API 教程 完成端到端训练,最后通过 KVStore 教程 将训练扩展到多 GPU 与多机集群。整套 API 建立在 src/operator 中 C++ 算子的高性能实现之上,这也是 MXNet 各语言绑定共享同一计算内核、模型格式互通互用的根本保证。

  • 人工智能
  • 深度学习
  • 机器学习

【免费下载链接】mxnet

Lightweight, 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
点击查看免费下载

相关推荐

上一篇:从阻塞到异步:Go语言优雅集成RabbitMQ的实战指南
下一篇:Mermaid图表实时编辑器:5个理由让你告别传统拖拽式图表工具

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询