
使用 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 维特征、绝大多数位置为 0RowSparseNDArray适合按行存储稀疏向量/矩阵常见于稀疏权重参数与稀疏梯度。二者的基础数据结构用法可分别参考仓库内的两份教程CSRNDArray - Compressed Sparse Row 存储格式教程 与 RowSparseNDArray - 稀疏梯度更新教程。在这两份教程之上MXNet 还提供了Sparse Symbol APImx.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是符号图中的占位符既可用于存放稠密数组也可用于存放稀疏数组。变量的存储类型 stypestypestorage 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, stypecsr) # 创建容纳 RowSparseNDArray 的变量 c mx.sym.Variable(c, styperow_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(ctxmx.cpu(), bshape) c_exec c.simple_bind(ctxmx.cpu(), cshape) 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.]], dtypefloat32)}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_sparseadd_exec mx.sym.Group([d, e, f]).simple_bind(ctxmx.cpu(), ashape, bshape, cshape) 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(ctxmx.cpu(), ashape, bshape, cshape) 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_VERBOSE0来抑制输出。检查符号图的存储类型当需要排查符号图中每个算子输入/输出的存储类型分配是否如预期时可设置环境变量MXNET_INFER_STORAGE_TYPE_VERBOSE_LOGGING1设置后MXNet 会把计算图中各算子输入输出的存储类型信息打印到控制台。对应实现在 src/executor/infer_graph_attr_pass.ccInferStorageType读取该环境变量为真时调用LogInferStorage输出整张图的存储类型与分发模式。例如检查一个典型的稀疏线性分类网络import mxnet as mx import os #os.environ[MXNET_INFER_STORAGE_TYPE_VERBOSE_LOGGING] 1 # 数据以 csr 格式输入 data mx.sym.var(data, stypecsr, shape(32, 10000)) # 权重以 row_sparse 格式存储 weight mx.sym.var(weight, styperow_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(datapred, labely, nameoutput) executor output.simple_bind(ctxmx.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)生成批大小设为 1last_batch_handlediscard丢弃末尾不完整批次label_namelabel与后续符号中的标签变量名保持一致。# 随机训练数据 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_handlediscard, label_namelabel)提示运行中出现的 SciPy 相关警告不影响本示例的正确性可忽略。定义模型模型的关键在于为每个变量声明合适的存储类型数据用csr可学习权重用row_sparse这样优化器会对该参数执行稀疏更新规则并用init属性指定该变量的初始化器initializer mx.initializer.Normal(sigma0.01) X mx.sym.Variable(data, stypecsr) Y mx.symbol.Variable(label) weight mx.symbol.Variable(weight, styperow_sparse, shape(feature_dimension, 1), initinitializer) bias mx.symbol.Variable(bias, shape(1, )) pred mx.sym.broadcast_add(mx.sym.sparse.dot(X, weight), bias) lro mx.sym.LinearRegressionOutput(datapred, labelY, namelro)该网络用到的符号及其作用如下Variable X稀疏数据输入占位符stypecsr声明其容纳 CSR 格式数组Variable Y稠密标签占位符Variable weight待学习权重styperow_sparse使其初始化为RowSparseNDArray且优化器将对其执行稀疏更新规则init指定该变量的初始化器Normal(sigma0.01)Variable bias待学习偏置sparse.dotX与weight的点积其稀疏实现会专门处理csr×row_sparse的组合见 src/operator/tensor/dot-inl.h 中针对 CSR 左乘 row_sparse/稠密右矩阵的分发逻辑输出存储类型推断为row_sparsebroadcast_add将bias广播加到点积结果上LinearRegressionOutput输出层计算输入与标签之间的 l2 损失。训练模型定义模型结构后创建 Module 并初始化参数与优化器# 创建 Module mod mx.mod.Module(symbollro, data_names[data], label_names[label]) # 依据迭代器提供的形状分配内存 mod.bind(data_shapestrain_iter.provide_data, label_shapestrain_iter.provide_label) # 用随机数初始化参数 mod.init_params(initializerinitializer) # 使用 SGD 优化器它对 row_sparse 权重执行稀疏更新 sgd mx.optimizer.SGD(learning_rate0.05, rescale_grad1.0/batch_size, momentum0.9) mod.init_optimizer(optimizersgd)要点说明data_names[data]与label_names[label]必须与符号中data、label变量的名字一一对应rescale_grad1.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_trainTrue) # 计算预测值 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_fnbatch_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 --launcherlocal python train.py --kvstoredist_async在单机启动 2 worker 2 server。模型定义 linear_model.py 与本教程的线性回归结构一致CSR 数据 × row_sparse 权重 →sparse.dot→broadcast_add→ 损失层。总结通过本教程你可以掌握 MXNet 稀疏符号的完整使用链路声明用stypecsr/styperow_sparse声明稀疏变量绑定与求值用simple_bind实例化执行器通过arg_dict更新变量持有的稀疏数组组合与推断用mx.sym.sparse算子组合符号图理解输出存储类型的自动推断对不支持稀疏的算子理解存储回退机制及其性能代价并通过MXNET_INFER_STORAGE_TYPE_VERBOSE_LOGGING1检查整图的存储类型分配训练用 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),仅供参考