news 2026/9/21 14:31:40

使用 MXNet Sparse Symbol 与 Module API 训练稀疏线性回归模型

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
使用 MXNet Sparse Symbol 与 Module API 训练稀疏线性回归模型

使用 MXNet Sparse Symbol 与 Module 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

本篇教程聚焦 MXNet 的Sparse Symbol API:如何声明可容纳稀疏数组的符号变量(csr/row_sparse存储类型)、理解符号图的存储类型推断与存储回退机制,并最终用 Module API 完成一个以 CSR 稀疏数据为输入、以 row_sparse 权重为可学习参数的线性回归模型的完整训练。读完本文,你将掌握稀疏符号的声明与绑定、mx.sym.sparse稀疏算子用法、存储类型推断与调试方法,以及利用稀疏梯度更新降低大模型通信开销的分布式训练要点。

背景:从稀疏 NDArray 到稀疏 Symbol

在 MXNet 中,CSRNDArray(压缩稀疏行格式)与RowSparseNDArray(行稀疏格式)是两种基本稀疏数据结构,用于高效表示零值占多数的张量:

  • CSRNDArray适合按行存储二维稀疏矩阵,常见于特征维度极高的训练样本(如 10000 维特征、绝大多数位置为 0);
  • RowSparseNDArray适合按行存储稀疏向量/矩阵,常见于稀疏权重参数与稀疏梯度。

二者的基础数据结构用法可分别参考仓库内的两份教程:CSRNDArray - Compressed Sparse Row 存储格式教程 与 RowSparseNDArray - 稀疏梯度更新教程。

在这两份教程之上,MXNet 还提供了Sparse Symbol APImx.sym.sparse包),让稀疏数组也能进入符号图(Symbolic Graph)的声明式表达:既可以作为图的输入占位符,也可以作为待学习的稀疏参数参与前向/反向计算。本教程将先用最小示例演示稀疏符号的基本用法,再完整训练一个线性回归模型。

前置条件

  • 已安装 MXNet(安装方式请参考仓库根目录 README.md 与 Setup and Installation 说明);
  • Python 环境,并安装jupyterrequests(教程以 Notebook 形式呈现时使用):
    pip install jupyter requests
  • 掌握 MXNet Symbol 的基本用法(变量、算子、自动微分),可参考仓库内 Symbol 相关文档与 python/mxnet/symbol 源码;
  • 掌握 CSRNDArray 与 RowSparseNDArray 的基础知识(见上文两份教程)。

稀疏符号变量(Variables)

变量(Variable)是符号图中的占位符,既可用于存放稠密数组,也可用于存放稀疏数组。

变量的存储类型 stype

stype(storage type)属性用于声明该变量所容纳数组的存储类型:

  • 默认值为"default",表示稠密存储格式;
  • 指定为"csr",表示将容纳CSRNDArray
  • 指定为"row_sparse",表示将容纳RowSparseNDArray
import mxnet as mx import numpy as np import random # 固定随机种子以保证结果可复现 random.seed(42) np.random.seed(42) mx.random.seed(42) # 创建容纳稠密 NDArray 的变量 a = mx.sym.Variable('a') # 创建容纳 CSRNDArray 的变量 b = mx.sym.Variable('b', stype='csr') # 创建容纳 RowSparseNDArray 的变量 c = mx.sym.Variable('c', stype='row_sparse') (a, b, c)

输出为:

(<Symbol a>, <Symbol b>, <Symbol c>)

可见stype只是变量的一种声明属性:bc在符号层面并未真正持有数据,真正决定数据形态的是绑定(bind)时喂入的数组。

用稀疏数组绑定变量

要计算一个稀疏符号,需要先实例化执行器(executor)。simple_bind会按各自由变量的存储类型分配全零数组作为初始值,然后通过forward方法求值,通过outputs属性取得全部输出。

shape = (2, 2) # 从稀疏符号实例化执行器 b_exec = b.simple_bind(ctx=mx.cpu(), b=shape) c_exec = c.simple_bind(ctx=mx.cpu(), c=shape) b_exec.forward() c_exec.forward() # b 与 c 被绑定为全零稀疏数组 print(b_exec.outputs, c_exec.outputs)

输出:

([ <CSRNDArray 2x2 @cpu(0)>], [ <RowSparseNDArray 2x2 @cpu(0)>])

绑定后可以通过执行器的arg_dict字典访问并更新变量所持有的数组。例如把b更新为全 1 的 CSR 数组:

b_exec.arg_dict['b'][:] = mx.nd.ones(shape).tostype('csr') b_exec.forward() # 变量 b 持有的数组已被更新为全 1 eval_b = b_exec.outputs[0] {'eval_b': eval_b, 'eval_b.asnumpy()': eval_b.asnumpy()}

输出:

{'eval_b': <CSRNDArray 2x2 @cpu(0)>, 'eval_b.asnumpy()': array([[ 1., 1.], [ 1., 1.]], dtype=float32)}

tostype('csr')是稠密数组到 CSR 稀疏数组的类型转换入口,其底层对应cast_storage算子(实现见 src/operator/tensor/cast_storage-inl.cuh),这也是后续训练数据准备阶段的核心转换手段。

符号组合与存储类型推断

稀疏算子与符号组合

稀疏符号可用算子组合成更复杂的表达式。mx.sym.sparse包提供专门的稀疏算子实现,例如对 CSR 输入取负的sparse.negative、针对 row_sparse 输入的元素级加法sparse.elemwise_add

# 稠密变量 a 的元素级加法(default stype) d = mx.sym.elemwise_add(a, a) # 对 csr 变量 b 取负 e = mx.sym.sparse.negative(b) # row_sparse 变量 c 的元素级加法 f = mx.sym.sparse.elemwise_add(c, c) {'d': d, 'e': e, 'f': f}

输出:

{'d': <Symbol elemwise_add0>, 'e': <Symbol negative0>, 'f': <Symbol elemwise_add1>}

从源码结构看,mx.sym.sparse(见 python/mxnet/symbol/sparse.py)是对底层稀疏算子的符号封装;这些算子的后端实现按照DispatchMode分发到不同的计算路径。以稀疏矩阵乘法为例,src/operator/tensor/dot-inl.h 中的存储类型推断逻辑会依据左/右输入是否为 CSR、是否转置等组合,把计算派发到kFCompute(默认实现)、kFComputeEx(稀疏专属实现)或kFComputeFallback(回退到稠密实现)。

存储类型推断

MXNet 中,任意稀疏符号的输出存储类型会根据输入存储类型自动推断。例如,elemwise_add(csr, csr)的输出会被推断为csrelemwise_add(row_sparse, row_sparse)的输出会被推断为row_sparse

add_exec = mx.sym.Group([d, e, f]).simple_bind(ctx=mx.cpu(), a=shape, b=shape, c=shape) add_exec.forward() dense_add = add_exec.outputs[0] # elemwise_add(csr, csr) 的输出存储类型推断为 "csr" csr_add = add_exec.outputs[1] # elemwise_add(row_sparse, row_sparse) 的输出存储类型推断为 "row_sparse" rsp_add = add_exec.outputs[2] {'dense_add.stype': dense_add.stype, 'csr_add.stype': csr_add.stype, 'rsp_add.stype': rsp_add.stype}

输出:

{'csr_add.stype': 'csr', 'dense_add.stype': 'default', 'rsp_add.stype': 'row_sparse'}

存储类型推断在 C++ 执行层由InferStorageType完成(见 src/executor/infer_graph_attr_pass.cc):它对图中每个节点调用算子的FInferStorageType特性,结合输入存储类型与算子自身的推断函数,得到所有节点的输出存储类型,并为每个节点标记dispatch_mode(默认实现 / 稀疏专属实现 / 回退实现),后续算子执行时据此选择具体内核。

存储类型回退(Storage Type Fallback)

并非所有算子都原生支持稀疏输入。对于不支持稀疏的稠密算子,MXNet 会自动执行存储类型回退

  • 若输入是稀疏数组,MXNet 会将其临时转换为稠密数组后调用稠密实现;
  • 若输出被指定为稀疏格式,MXNet 会把稠密算子的输出再转换回目标稀疏格式

整个过程不影响计算正确性,但会带来一定的性能开销,且发生回退时控制台会打印警告信息。

# `log` 算子完全不支持稀疏输入,可回退到稠密实现 csr_log = mx.sym.log(a) # `elemwise_add` 不支持 csr 与 row_sparse 混合相加,可回退到稠密实现 csr_rsp_add = mx.sym.elemwise_add(b, c) fallback_exec = mx.sym.Group([csr_rsp_add, csr_log]).simple_bind(ctx=mx.cpu(), a=shape, b=shape, c=shape) fallback_exec.forward() fallback_add = fallback_exec.outputs[0] fallback_log = fallback_exec.outputs[1] {'fallback_add': fallback_add, 'fallback_log': fallback_log}

输出(两个输出均为稠密 NDArray):

{'fallback_add': [[ 0. 0.] [ 0. 0.]] <NDArray 2x2 @cpu(0)>, 'fallback_log': [[-inf -inf] [-inf -inf]] <NDArray 2x2 @cpu(0)>}

回退警告的实现在 src/common/utils.h:一旦算子因存储类型不匹配而走默认(稠密)实现,LogStorageFallback会打印诸如"The operator with default storage type will be dispatched for execution... Temporary dense ndarrays are generated in order to execute the operator"的提示。若确定回退不影响业务且不想看到警告,可设置环境变量MXNET_STORAGE_FALLBACK_LOG_VERBOSE=0来抑制输出。

检查符号图的存储类型

当需要排查符号图中每个算子输入/输出的存储类型分配是否如预期时,可设置环境变量:

MXNET_INFER_STORAGE_TYPE_VERBOSE_LOGGING=1

设置后,MXNet 会把计算图中各算子输入输出的存储类型信息打印到控制台。对应实现在 src/executor/infer_graph_attr_pass.cc:InferStorageType读取该环境变量,为真时调用LogInferStorage输出整张图的存储类型与分发模式。

例如,检查一个典型的稀疏线性分类网络:

import mxnet as mx import os #os.environ['MXNET_INFER_STORAGE_TYPE_VERBOSE_LOGGING'] = "1" # 数据以 csr 格式输入 data = mx.sym.var('data', stype='csr', shape=(32, 10000)) # 权重以 row_sparse 格式存储 weight = mx.sym.var('weight', stype='row_sparse', shape=(10000, 2)) bias = mx.symbol.Variable("bias", shape=(2,)) dot = mx.symbol.sparse.dot(data, weight) pred = mx.symbol.broadcast_add(dot, bias) y = mx.symbol.Variable("label") output = mx.symbol.SoftmaxOutput(data=pred, label=y, name="output") executor = output.simple_bind(ctx=mx.cpu())

取消os.environ[...]一行的注释后运行,即可在控制台看到sparse.dot(csr, row_sparse)的输出被推断为row_sparsebroadcast_addSoftmaxOutput的输入输出存储类型等关键信息,从而验证整条链路的稀疏性是否按预期传递。

使用 Module API 训练稀疏线性回归

下面用稀疏符号 + 稀疏优化器完整实现一个线性回归模型。待拟合的目标函数为:

y = x1 + 2·x2 + 3·x3 + ... + 100·x100

其中(x1, x2, ..., x100)为输入特征,y为对应标签。

准备数据

mx.io.LibSVMItermx.io.NDArrayIter都支持以 CSR 格式加载稀疏数据。本例使用NDArrayIter

  • mx.test_utils.rand_ndarray生成一个 1000×100 的 CSR 稀疏训练矩阵(密度 0.01,即只有约 1% 的元素非零),其实现见 python/mxnet/test_utils.py;
  • 真实权重为1..100,标签通过稀疏矩阵乘法mx.nd.dot(train_data, target_weight)生成;
  • 批大小设为 1,last_batch_handle='discard'丢弃末尾不完整批次,label_name='label'与后续符号中的标签变量名保持一致。
# 随机训练数据 feature_dimension = 100 train_data = mx.test_utils.rand_ndarray((1000, feature_dimension), 'csr', 0.01) target_weight = mx.nd.arange(1, feature_dimension + 1).reshape((feature_dimension, 1)) train_label = mx.nd.dot(train_data, target_weight) batch_size = 1 train_iter = mx.io.NDArrayIter(train_data, train_label, batch_size, last_batch_handle='discard', label_name='label')

提示:运行中出现的 SciPy 相关警告不影响本示例的正确性,可忽略。

定义模型

模型的关键在于为每个变量声明合适的存储类型:数据用csr,可学习权重用row_sparse(这样优化器会对该参数执行稀疏更新规则),并用init属性指定该变量的初始化器:

initializer = mx.initializer.Normal(sigma=0.01) X = mx.sym.Variable('data', stype='csr') Y = mx.symbol.Variable('label') weight = mx.symbol.Variable('weight', stype='row_sparse', shape=(feature_dimension, 1), init=initializer) bias = mx.symbol.Variable('bias', shape=(1, )) pred = mx.sym.broadcast_add(mx.sym.sparse.dot(X, weight), bias) lro = mx.sym.LinearRegressionOutput(data=pred, label=Y, name="lro")

该网络用到的符号及其作用如下:

  1. Variable X:稀疏数据输入占位符,stype='csr'声明其容纳 CSR 格式数组;
  2. Variable Y:稠密标签占位符;
  3. Variable weight:待学习权重,stype='row_sparse'使其初始化为RowSparseNDArray,且优化器将对其执行稀疏更新规则;init指定该变量的初始化器(Normal(sigma=0.01));
  4. Variable bias:待学习偏置;
  5. sparse.dotXweight的点积,其稀疏实现会专门处理csr×row_sparse的组合(见 src/operator/tensor/dot-inl.h 中针对 CSR 左乘 row_sparse/稠密右矩阵的分发逻辑),输出存储类型推断为row_sparse
  6. broadcast_add:将bias广播加到点积结果上;
  7. LinearRegressionOutput:输出层,计算输入与标签之间的 l2 损失。

训练模型

定义模型结构后,创建 Module 并初始化参数与优化器:

# 创建 Module mod = mx.mod.Module(symbol=lro, data_names=['data'], label_names=['label']) # 依据迭代器提供的形状分配内存 mod.bind(data_shapes=train_iter.provide_data, label_shapes=train_iter.provide_label) # 用随机数初始化参数 mod.init_params(initializer=initializer) # 使用 SGD 优化器,它对 "row_sparse" 权重执行稀疏更新 sgd = mx.optimizer.SGD(learning_rate=0.05, rescale_grad=1.0/batch_size, momentum=0.9) mod.init_optimizer(optimizer=sgd)

要点说明:

  • data_names=['data']label_names=['label']必须与符号中datalabel变量的名字一一对应;
  • rescale_grad=1.0/batch_size将梯度按批大小归一化,与优化器学习率配合得到稳定的更新步长;
  • 使用SGD 作为稀疏优化器:对row_sparse参数,优化器只更新梯度中非零行对应的权重行,这正是稀疏更新规则节省计算与通信的核心。

最后用 Module 的forward/backward/update三个方法驱动训练循环,并以 MSE(均方误差)作为评估指标:

# 使用均方误差作为评估指标 metric = mx.metric.create('MSE') # 训练 10 个 epoch for epoch in range(10): train_iter.reset() metric.reset() for batch in train_iter: mod.forward(batch, is_train=True) # 计算预测值 mod.update_metric(metric, batch.label) # 累积预测误差指标 mod.backward() # 计算梯度 mod.update() # 更新参数 print('Epoch %d, Metric = %s' % (epoch, metric.get())) assert metric.get()[1] < 1, "Achieved MSE (%f) is larger than expected (1.0)" % metric.get()[1]

训练 10 个 epoch 后 MSE 收敛到 1 以下(示例输出:Epoch 9, Metric = ('mse', 0.35979430613957991)),assert确保最终 MSE 小于阈值 1.0,即模型成功拟合了目标函数。

多机 / 多设备分布式训练

MXNet 支持对row_sparse权重和梯度进行分布式训练,可显著降低大模型的通信开销。分布式训练时需要注意:当使用 KVStore 在多个设备/机器间更新参数时,update()只更新 KVStore 中的参数副本,并不会自动把更新后的参数广播回所有设备。因此需要调用prepare来按下一批数据的行索引拉取稀疏权重,即mod.prepare(batch, sparse_row_id_fn=...),在forwardsave_checkpoint之前必须完成该步骤。

prepare的语义在 python/mxnet/module/module.py 中有明确定义:sparse_row_id_fn是一个回调函数,接收data_batch并返回{参数名: 行ID数组}的字典,Module 据此从 KVStore 中拉取row_sparse参数对应行的最新值。仓库内的完整可运行示例见 example/sparse/linear_classification,其核心流程如下:

  • 使用mx.io.LibSVMIter加载 LibSVM 格式的稀疏数据(如 Avazu 点击率预测数据集,特征维度高达 100 万);
  • 定义batch_row_ids(返回当前 mini-batch 非零行索引data_batch.data[0].indices)与all_row_ids(返回全部行索引)两个行 ID 回调;
  • 训练循环中先mod.prepare(batch, sparse_row_id_fn=batch_row_ids)mod.forward_backward(batch)mod.update()
  • 评估与保存 checkpoint 前调用mod.prepare(None, all_row_ids)拉取全部权重行。

相关脚本 example/sparse/linear_classification/train.py 还演示了通过--kvstore参数在dist_syncdist_asynclocal三种模式间切换,以及配合 tools/launch.py 启动多 worker 多 server 集群的方式(例如python ../../../tools/launch.py -n 2 --launcher=local python train.py --kvstore=dist_async在单机启动 2 worker + 2 server)。模型定义 linear_model.py 与本教程的线性回归结构一致:CSR 数据 × row_sparse 权重 →sparse.dotbroadcast_add→ 损失层。

总结

通过本教程,你可以掌握 MXNet 稀疏符号的完整使用链路:

  1. 声明:用stype='csr'/stype='row_sparse'声明稀疏变量;
  2. 绑定与求值:用simple_bind实例化执行器,通过arg_dict更新变量持有的稀疏数组;
  3. 组合与推断:用mx.sym.sparse算子组合符号图,理解输出存储类型的自动推断;对不支持稀疏的算子,理解存储回退机制及其性能代价,并通过MXNET_INFER_STORAGE_TYPE_VERBOSE_LOGGING=1检查整图的存储类型分配;
  4. 训练:用 Module API + 稀疏优化器(SGD)训练row_sparse权重,数据以 CSR 格式经NDArrayIter/LibSVMIter喂入;
  5. 扩展:通过mod.prepare(batch, sparse_row_id_fn=...)支持多机/多设备分布式稀疏训练,参考 example/sparse/linear_classification 与 example/sparse 下的更多稀疏示例(矩阵分解、因子分解机、wide & deep 等)继续深入。

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

Qt离线安装全攻略:从选型到Kit配置的完整指南

1. 为什么离线装 Qt 这件事值得单独写一篇如果你所在的项目环境是内网、工控机、涉密终端&#xff0c;或者客户现场压根没有外网&#xff0c;那你迟早会撞上“Qt 离线安装”这堵墙。在线安装器走不通&#xff0c;apt、yum、pip全部失效&#xff0c;连下载一个 30MB 的 MinGW 都…

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

Ubuntu 20.04 离线安装 Realtek RTL8852BE 无线网卡驱动实战

装过 Linux 的朋友基本都有类似遭遇&#xff1a;系统装好了&#xff0c;界面也正常&#xff0c;结果右上角偏偏没有 WiFi 图标。尤其是一台崭新的笔记本&#xff0c;或者刚换的 USB 无线网卡&#xff0c;插上去一点反应没有&#xff0c;那一刻的心情真的有点崩溃。这次要聊的就…

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

Flutter图标颜色在鸿蒙系统的适配方案

1. 项目背景与核心挑战在跨平台开发领域&#xff0c;Flutter框架因其高效的渲染性能和丰富的组件库而广受欢迎。而鸿蒙系统作为新兴的操作系统平台&#xff0c;其设计理念和实现机制与传统Android/iOS存在显著差异。当开发者尝试将现有Flutter应用迁移到鸿蒙平台时&#xff0c;…

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

Mirror网络库自定义生成函数实战指南

1. Mirror网络库自定义生成函数深度解析在多人联机游戏开发中&#xff0c;对象生成与销毁是最基础也最关键的环节之一。Mirror作为Unity的高性能网络库&#xff0c;默认提供了简单的预制体实例化机制&#xff0c;但在实际项目中&#xff0c;我们往往需要更精细的控制——比如对…

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

VibeCoding 做历史粘贴板,Claude Code 的模型通道走 TaoToken

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

作者头像 李华