使用 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 API(mx.sym.sparse包),让稀疏数组也能进入符号图(Symbolic Graph)的声明式表达:既可以作为图的输入占位符,也可以作为待学习的稀疏参数参与前向/反向计算。本教程将先用最小示例演示稀疏符号的基本用法,再完整训练一个线性回归模型。
前置条件
- 已安装 MXNet(安装方式请参考仓库根目录 README.md 与 Setup and Installation 说明);
- Python 环境,并安装
jupyter与requests(教程以 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只是变量的一种声明属性:b与c在符号层面并未真正持有数据,真正决定数据形态的是绑定(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)的输出会被推断为csr,elemwise_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_sparse、broadcast_add与SoftmaxOutput的输入输出存储类型等关键信息,从而验证整条链路的稀疏性是否按预期传递。
使用 Module API 训练稀疏线性回归
下面用稀疏符号 + 稀疏优化器完整实现一个线性回归模型。待拟合的目标函数为:
y = x1 + 2·x2 + 3·x3 + ... + 100·x100其中(x1, x2, ..., x100)为输入特征,y为对应标签。
准备数据
mx.io.LibSVMIter与mx.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")该网络用到的符号及其作用如下:
Variable X:稀疏数据输入占位符,stype='csr'声明其容纳 CSR 格式数组;Variable Y:稠密标签占位符;Variable weight:待学习权重,stype='row_sparse'使其初始化为RowSparseNDArray,且优化器将对其执行稀疏更新规则;init指定该变量的初始化器(Normal(sigma=0.01));Variable bias:待学习偏置;sparse.dot:X与weight的点积,其稀疏实现会专门处理csr×row_sparse的组合(见 src/operator/tensor/dot-inl.h 中针对 CSR 左乘 row_sparse/稠密右矩阵的分发逻辑),输出存储类型推断为row_sparse;broadcast_add:将bias广播加到点积结果上;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']必须与符号中data、label变量的名字一一对应;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=...),在forward或save_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_sync、dist_async、local三种模式间切换,以及配合 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.dot→broadcast_add→ 损失层。
总结
通过本教程,你可以掌握 MXNet 稀疏符号的完整使用链路:
- 声明:用
stype='csr'/stype='row_sparse'声明稀疏变量; - 绑定与求值:用
simple_bind实例化执行器,通过arg_dict更新变量持有的稀疏数组; - 组合与推断:用
mx.sym.sparse算子组合符号图,理解输出存储类型的自动推断;对不支持稀疏的算子,理解存储回退机制及其性能代价,并通过MXNET_INFER_STORAGE_TYPE_VERBOSE_LOGGING=1检查整图的存储类型分配; - 训练:用 Module API + 稀疏优化器(SGD)训练
row_sparse权重,数据以 CSR 格式经NDArrayIter/LibSVMIter喂入; - 扩展:通过
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),仅供参考