MXNet NDArray 完全指南:多维数组创建、运算、上下文管理与稀疏存储实战
【免费下载链接】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
NDArray 是 Apache MXNet(本仓库 mxnet)中承载数据与模型参数的核心多维数组类型,几乎所有训练与推理流程都围绕它展开。本文以 docs/python_docs/python/tutorials/packages/legacy/ndarray/ 下的系列教程为主体,系统讲解 NDArray 的创建方式、数学运算、原地操作、切片与广播、CPU/GPU 上下文管理,并深入剖析 CSRNDArray 与 RowSparseNDArray 两种稀疏存储格式及其在稀疏梯度更新中的应用。读完本文,你将能够用 NDArray 完成从数据准备到多设备训练的全部基础操作,并理解其异步执行引擎带来的性能特性。
1. NDArray 是什么:MXNet 的数据心脏
在 MXNet 中,NDArray是一个表示多维、同构、固定大小元素数组的对象,其 Python 侧类型定义位于 python/mxnet/ndarray/ndarray.py#L266-L288。它承担两个核心角色:
- 数据容器:存放训练样本、中间特征、梯度等一切张量数据;
- 模型参数载体:神经网络权重与偏置都以 NDArray 形式存在,并可通过
.context属性被放置在 CPU 或指定 GPU 上。
NDArray API 的设计刻意与 NumPy 保持高度相似,方便数据科学家快速上手;但它与 NumPy 有一个本质区别——计算是异步、非阻塞的。当执行c = a * b(其中a、b均为 NDArray)时,操作被提交给 MXNet 的 Execution Engine(执行引擎),函数立即返回,用户线程可继续执行后续代码,即使计算尚未完成。引擎会构建计算图,在保证依赖顺序的前提下对计算进行重排或合并:若后续代码对c有操作,引擎会在c的结果就绪后自动启动它们,无需手写回调。Engine 的异步调度实现可参考 src/engine/threaded_engine.cc#L317 的PushAsync等接口。
需要强调:自 MXNet 1.x 起,MXNet 提供了 NumPy 兼容的新数组类mx.np.ndarray,本文所述的经典NDArray(即mx.nd)属于 legacy(遗留)数组类,相关教程目录见 docs/python_docs/python/tutorials/packages/legacy/ndarray/index.rst。
2. 创建 NDArray 的七种姿势
2.1 从 Python 列表创建:nd.array
最简单的方式是直接由 Python 列表构造一维或二维数组:
import mxnet as mx from mxnet import nd x = nd.array([1, 2, 3]) # 一维数组 print(x) y = nd.array([[1, 2, 3, 4], [1, 2, 3, 4], [1, 2, 3, 4]]) # 3x4 二维矩阵 print(y)2.2 未初始化与填充创建:empty/full
nd.empty(shape):只申请内存、不初始化元素值,返回的矩阵内容不可预测(可能包含任意大数);nd.full(shape, value):用指定值填充整个数组。
x = nd.empty((3, 3)) print(x) x = nd.full((3, 3), 7) print(x)2.3 全零与全一:zeros/ones
绝大多数场景我们希望数组被确定性初始化,最常用的是全零矩阵:
x = nd.zeros((3, 10)) print(x)全一矩阵对应nd.ones,例如nd.ones((3, 4))会生成一个 3 行 4 列、元素全为 1 的矩阵。
2.4 随机初始化:random_normal
神经网络参数初始化时常需要随机采样。nd.random_normal(loc, scale, shape=(...))从均值为loc、方差为scale的正态分布中采样:
y = nd.random_normal(0, 1, shape=(3, 4)) print(y)2.5 按形状复制:zeros_like
zeros_like复制一个数组的形状但内容置零,常用于准备与某个张量同形的输出缓冲区:
z = nd.zeros_like(y) print(z)2.6 等间隔序列:arange
nd.arange(n)生成[0, 1, ..., n-1]的等间隔一维数组,配合.reshape可一步生成自定义形状数据:
y = nd.arange(6) y = y.reshape((3, 1)) # 或直接链式调用 y = nd.arange(6).reshape((3, 1))2.7 查询数组属性:shape/size/dtype/context
每个 NDArray 都可通过属性查询其元信息:
y.shape # 各维度大小组成的元组 y.size # 元素总数,等于 shape 各分量之积 y.dtype # 数据类型 y.context # 数据所在设备(CPU 或某块 GPU)float32是默认数据类型。可以显式指定其他精度——低精度可提升性能,高精度保证数值稳定性:
import numpy as np a = nd.array([1, 2, 3]) # 默认 float32 b = nd.array([1, 2, 3], dtype=np.int32) # 32 位整型 c = nd.array([1.2, 2.3], dtype=np.float16) # 16 位半精度浮点 (a.dtype, b.dtype, c.dtype)完整创建 API 的文档化说明可对照 python/mxnet/ndarray/ndarray.py 中对应工厂函数。
3. 数组运算:从逐元素到矩阵乘法
NDArray 支持大量标准数学运算,且运算符被重载,写起来与 NumPy 几乎一致。__add__、__mul__等运算符重载的实现可见 python/mxnet/ndarray/ndarray.py#L322-L379。
3.1 逐元素运算
x = nd.ones((3, 4)) y = nd.random_normal(0, 1, shape=(3, 4)) print('x=', x) print('y=', y) x = x + y # 逐元素加法 print('x = x + y, x=', x) x = nd.array([1, 2, 3]) y = nd.array([2, 2, 2]) print(x * y) # 逐元素乘法 print(nd.exp(x)) # 逐元素指数运算3.2 转置与矩阵乘法
对于二维矩阵,先转置再点积即可完成真正的矩阵乘法:
nd.dot(x, y.T)nd.dot是 MXNet 中最常用的矩阵运算算子之一,其底层在 CPU/GPU 上均有高度优化实现。
4. 原地操作与内存管理
4.1 每次运算都会分配新内存
每次执行y = x + y都会分配一块新内存存放结果,然后让y指向新对象,旧内存被释放。用 Python 内置的id()函数可以验证变量引用的对象是否改变:
print('y=', y) print('id(y):', id(y)) y = y + x print('after y=y+x, y=', y) print('id(y):', id(y)) # id 已经改变4.2 切片写入复用缓冲区:result[:] = ...
若想复用已分配的内存,可用切片赋值语法:
z = nd.zeros_like(x) print('id(z):', id(z)) z[:] = x + y print('z[:] = x + y, z=', z) print('id(z) is the same as before:', id(z)) # id 保持不变不过x + y这一步仍会分配一个临时缓冲区,先把结果算出来再拷入z。
4.3out=参数:彻底消除临时缓冲区
每个算子都支持out关键字参数,直接把结果写入指定数组:
nd.elemwise_add(x, y, out=z) print('after nd.elemwise_add(x, y, out=z), z=', z, 'is in id(z):', id(z))z的id始终不变,全程无临时分配。__iadd__等原地运算符(x += y)内部正是通过op.broadcast_add(self, other, out=self)实现的(见 python/mxnet/ndarray/ndarray.py#L326-L335),效果与显式out=相同。
4.4 两种原地写法
不打算复用x时,可把结果写回x本身,MXNet 提供两种方式:
# 方式一:切片赋值 x[:] = x + y # 方式二:op-equals 运算符 x += y在+=这类原地操作中,MXNet 还会校验数组是否可写(writable),只读数组会抛出ValueError。
5. 切片与广播
5.1 切片语法速查
NDArray 完整支持 NumPy 风格的切片:
a[start:end]:取start到end-1的元素;a[start:]:取start到末尾;a[:end]:取开头到end-1;a[:]:整个数组的副本。
一维与二维读取示例:
x = nd.array([1, 2, 3]) s = x[1:3] # 取第 2、3 个元素 print('slicing the 2nd and 3rd elements, s=', s) x = nd.array([[1, 2, 3, 4], [5, 6, 7, 8], [9, 10, 11, 12]]) s = x[1:3] # 取第 2、3 行 print('slicing the 2nd and 3rd rows, s=', s)5.2 写入与多维切片
切片不仅能读,还能写:
x[2] = 9.0 # 整行替换 x[0, 2] = 9.0 # 替换单个元素 x[1:2, 1:3] = 5.0 # 替换一块子区域多维切片同样支持按行列抽取:
x = nd.array([[1, 2, 3, 4], [5, 6, 7, 8], [9, 10, 11, 12]]) s = x[1:2, 1:3] # 取第 2 行第 2~3 列 s = x[:, :1] # 第一列 s = x[:1, :] # 第一行 s = x[:, 3:] # 最后一列 s = x[2:, :] # 最后一行5.3 广播(Broadcasting)
当低维数组与高维数组做逐元素运算时,MXNet 会触发广播机制:低维数组沿维度为 1 的轴复制扩展,直至形状匹配。
x = nd.ones(shape=(3, 6)) y = nd.arange(6) print('x + y = ', x + y)初始时y的形状是(6),MXNet 会将其推断为(1, 6),然后沿行方向广播成(3, 6)再相加。广播优先沿最左侧的轴复制,因此y被解释为(1, 6)而非(6, 1)。若想按列广播,需用.reshape显式给出二维形状:
y = y.reshape((3, 1)) print('x + y = ', x + y) # 按列广播 y = nd.arange(6).reshape((3, 1)) # arange 与 reshape 一步链式完成6. 上下文管理:CPU 与 GPU
6.1 每个数组都有上下文
MXNet 中每个数组都有一个 context(上下文),可以是 CPU,也可以是某块 GPU,甚至分布式场景下的多台服务器。合理分配数组所在设备,能最小化设备间数据传输时间——例如在带 GPU 的服务器上训练时,模型参数最好常驻 GPU。
将数组创建在指定设备上,通过ctx参数:
from mxnet import gpu from mxnet import nd z = nd.ones(shape=(3, 3), ctx=gpu(0)) # 放在第一块 GPU 上 print(z)没有 GPU 时,把ctx=gpu(0)替换为ctx=mx.cpu()即可。
6.2 跨设备复制:copyto
copyto把一个数组复制到目标上下文:
x_gpu = x.copyto(gpu(0)) print(x_gpu)运算符的结果与输入处于同一上下文:
x_gpu + z # 结果仍在 GPU(0) 上6.3 条件复制:as_in_context与copyto的区别
注意:即使数组已经在目标设备上,z.copyto(gpu(0))仍会复制一份并分配新内存。如果只是想“确保数组在正确设备上”,应使用as_in_context()——当变量已经在目标设备时,它是一次 no-op(空操作),不产生复制:
print('id(z):', id(z)) z = z.copyto(gpu(0)) # 总是分配新内存,id 改变 print('id(z):', id(z)) z = z.as_in_context(gpu(0)) # 已在 gpu(0) 上,no-op,id 不变 print('id(z):', id(z)) print(z)as_in_context与copyto的实现可参考 python/mxnet/ndarray/ndarray.py#L2716(copyto)与 python/mxnet/ndarray/ndarray.py#L2862(as_in_context)。训练循环中,建议始终用as_in_context处理数据批次与参数,避免无谓的跨设备拷贝。
7. 与 NumPy 互转及性能陷阱
7.1 互转接口
NDArray 与 NumPy 互转非常方便,且转换后的数组不共享内存:
a = x.asnumpy() # NDArray -> numpy.ndarray type(a) y = nd.array(a) # numpy.ndarray -> NDArray7.2 阻塞调用会打断异步流水线
.asnumpy()、.asscalar()、.wait_to_read()、.waitall()都是阻塞操作:调用时 MXNet 必须等待 Execution Engine 完成此前提交的所有异步计算,才能取回结果(.asnumpy()的 C++ 侧同步实现即MXNDArraySyncCopyToCPU,见 python/mxnet/ndarray/ndarray.py#L2652-L2659;.asscalar()位于 python/mxnet/ndarray/ndarray.py#L2661)。
这带来的实际体验是:如果计算图很长,某处调用.asnumpy()时会“感觉”耗时很久——真正耗时的并非转换本身,而是引擎要在此刻集中完成积压的全部异步计算。在 GPU 上这个问题更突出:数据必须从 GPU 拷贝回 CPU 才能生成np.array。
7.3 用 NDArray 原生算子替代 NumPy
性能最佳实践是直接在 NDArray 上进行所有运算,完全绕开 NumPy。当某个 NumPy 算子缺失时,可采取三种策略:
策略一:用若干低层算子组合出高层算子。例如 NumPy 有np.full_like,NDArray API 没有,但可以用ones乘以填充值替代:
from mxnet import nd import numpy as np np_y = np.full_like(a=np.arange(6, dtype=int), fill_value=10) nd_y = nd.ones(shape=(6,)) * 10 np.array_equal(np_y, nd_y.asnumpy()) # True策略二:寻找名称或签名相近的算子。例如nd.ravel_multi_index类似np.ravel;np.split按索引切分,而nd.split需要传入切分数。再如nd.pad只能处理 4 维张量,低维输入需先扩维再还原:
def pad_array(data, max_length): # 扩展到 4 维,因为 nd.pad 只支持 4 维张量 data_expanded = data.reshape(1, 1, 1, data.shape[0]) # 用常量 0 填充全部 4 个维度 data_padded = nd.pad(data_expanded, mode='constant', pad_width=[0, 0, 0, 0, 0, 0, 0, max_length - data.shape[0]], constant_value=0) # 移除临时维度 data_reshaped_back = data_padded.reshape(max_length) return data_reshaped_back pad_array(nd.array([1, 2, 3]), max_length=10) # 输出: [ 1. 2. 3. 0. 0. 0. 0. 0. 0. 0.] # <NDArray 10 @cpu(0)>7.4 最小化阻塞影响:延迟取值的 LossBuffer 模式
当不得不使用.asnumpy()或.asscalar()(例如打印 loss 指标)时,尽量在“该值大概率已算完”的时刻再取值。经典做法是用一个缓冲类缓存上一轮的 loss,把打印推迟一个迭代,让引擎有充足时间并行完成上一轮计算:
from __future__ import print_function import mxnet as mx from mxnet import gluon, nd, autograd from mxnet.ndarray import NDArray from mxnet.gluon import HybridBlock import numpy as np class LossBuffer(object): """存储 loss 值的简单缓冲,new_loss 返回上一轮 loss""" def __init__(self): self._loss = None def new_loss(self, loss): ret = self._loss self._loss = loss return ret @property def loss(self): return self._loss net = gluon.nn.Dense(10) ce = gluon.loss.SoftmaxCELoss() net.initialize() data = nd.random.uniform(shape=(1024, 100)) label = nd.array(np.random.randint(0, 10, (1024,)), dtype='int32') train_dataset = gluon.data.ArrayDataset(data, label) train_data = gluon.data.DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=2) trainer = gluon.Trainer(net.collect_params(), optimizer='sgd') loss_buffer = LossBuffer() for data, label in train_data: with autograd.record(): out = net(data) # 保存新 loss,返回上一轮 loss prev_loss = loss_buffer.new_loss(ce(out, label)) loss_buffer.loss.backward() trainer.step(data.shape[0]) if prev_loss is not None: print("Loss: {}".format(np.mean(prev_loss.asnumpy())))运行时会观察到 loss 输出大约如下(每次迭代延迟一个周期打印,阻塞时间被摊薄):
Loss: 2.310760974884033 Loss: 2.334498643875122 Loss: 2.3244147300720215 ...8. 稀疏 NDArray(一):CSRNDArray 压缩稀疏行格式
8.1 为什么需要稀疏存储
现实世界大量数据集是高维稀疏的。以推荐系统为例,类别与用户数量可达百万量级,但每个用户实际购买的类目极少,绝大多数元素为 0。用默认的稠密结构存储这类矩阵,内存与计算都被浪费在 0 上。
CSRNDArray以压缩稀疏行(CSR)格式存储二维矩阵,并让算子使用专门算法。该格式面向列数众多、且每一行都很稀疏(非零元素少)的 2D 矩阵。对高稀疏度矩阵(如约 1% 非零、密度约 1%),相比稠密NDArray有两个主要优势:
- 内存占用显著降低;
- 部分运算(如矩阵-向量乘法)明显更快。
CSR 格式与 SciPy 的实现相似,但CSRNDArray额外继承了 NDArray 的异步非阻塞求值与自动并行化能力。同时,NDArray家族新增了stype属性用于记录存储类型:稠密 NDArray 的stype为"default",CSRNDArray为"csr"。
8.2 CSR 三数组表示
一个 CSRNDArray 把二维矩阵表示为三个一维数组data、indptr、indices:
- data:矩阵非零元素(按行主序排列);
- indices:
data中每个非零元素对应的列索引; - indptr:每行第一个非零元素在
data中的偏移指针,indptr[0]恒为 0,indptr[i+1]是到第i行为止非零元素总数的累计值。
行i的列索引位于indices[indptr[i]:indptr[i+1]],对应值位于data[indptr[i]:indptr[i+1]]。同一行内列索引必须升序排列,且不允许出现重复列索引。
以矩阵为例:
[[7, 0, 8, 0] [0, 0, 0, 0] [0, 9, 0, 0]]按行主序去除全部 0 得到data = [7, 8, 9];每个非零元素所在列依次为第 0、2、1 列,即indices = [0, 2, 1];第一行有 2 个非零元素、第二行 0 个、第三行 1 个,累计偏移得到indptr = [0, 2, 2, 3]。重建时用data[0:2]+indices[0:2]还原第一行,data[2:2]+indices[2:2]还原全零行,data[2:3]+indices[2:3]还原第三行。
8.3 创建 CSRNDArray
方式一:由data/indices/indptr三元组创建(可传 Python 列表或 NumPy 数组):
import mxnet as mx shape = (3, 4) data_list = [7, 8, 9] indices_list = [0, 2, 1] indptr_list = [0, 2, 2, 3] a = mx.nd.sparse.csr_matrix((data_list, indices_list, indptr_list), shape=shape) a.asnumpy() # array([[ 7., 0., 8., 0.], # [ 0., 0., 0., 0.], # [ 0., 9., 0., 0.]], dtype=float32)mx.nd.sparse.csr_matrix的完整签名(含shape/ctx/dtype参数及其默认行为)定义在 python/mxnet/ndarray/sparse.py#L838:默认上下文为当前默认上下文,默认 dtype 为float32(当输入为 NDArray/NumPy 数组时沿用其 dtype)。
方式二:从 SciPy CSR 对象创建:
import numpy as np import scipy.sparse as spsp data_np = np.array([7, 8, 9]) indptr_np = np.array([0, 2, 2, 3]) indices_np = np.array([0, 2, 1]) c = spsp.csr.csr_matrix((data_np, indices_np, indptr_np), shape=shape) d = mx.nd.sparse.array(c) # scipy csr -> CSRNDArray print(d.asnumpy())方式三:从稠密数组压缩转换。已有数据但未计算indices/indptr时,可用tostype('csr')一步完成压缩,并直接访问压缩后的内部数组:
big_array = mx.nd.round(mx.nd.random.uniform(low=0, high=1, shape=(1000, 100))) big_array_csr = big_array.tostype('csr') indices = big_array_csr.indices indptr = big_array_csr.indptr data = big_array_csr.data # data + indices + indptr 的总大小远小于稠密 big_array指定元素类型:mx.nd.sparse.array(a)默认float32,也可通过dtype指定,例如mx.nd.array(a, dtype=np.float16)生成半精度数组。
8.4 检查与存储类型转换
检查 CSR 数组的常用方法:
.asnumpy():转成稠密numpy.ndarray查看内容;.data/.indices/.indptr:查看内部三个存储数组;.stype:查看存储类型('csr')。
{'a.stype': a.stype, 'data': a.data, 'indices': a.indices, 'indptr': a.indptr} # {'a.stype': 'csr', 'data': [ 7. 8. 9.] <NDArray 3 @cpu(0)>, # 'indices': [0 2 1] <NDArray 3 @cpu(0)>, 'indptr': [0 2 2 3] <NDArray 4 @cpu(0)>}存储类型转换有两种等价途径:
# 途径一:tostype 方法 ones = mx.nd.ones((2, 2)) csr = ones.tostype('csr') # default -> csr dense = csr.tostype('default') # csr -> default # 途径二:cast_storage 算子 csr = mx.nd.sparse.cast_storage(ones, 'csr') dense = mx.nd.sparse.cast_storage(csr, 'default')8.5 复制与索引
copy():深拷贝,返回新数组;copyto(dest)或切片赋值dest[:] = src:深拷贝到既有数组;- 注意:若源与目标存储类型不一致,
copyto/[]不会改变目标数组的存储类型(源数组会被临时转换)。
a = mx.nd.ones((2, 2)).tostype('csr') b = a.copy() c = mx.nd.sparse.zeros('csr', (2, 2)) c[:] = a d = mx.nd.sparse.zeros('csr', (2, 2)) a.copyto(d) # b/c/d 内容均为全 1,且 b is a 为 FalseCSRNDArray 支持沿 axis 0 切片,切片返回新 CSRNDArray;多维索引或沿特定轴的切片目前不支持:
a = mx.nd.array(np.arange(6).reshape(3, 2)).tostype('csr') b = a[1:2].asnumpy() # array([[ 2., 3.]]) c = a[:].asnumpy() # array([[ 0., 1.], [ 2., 3.], [ 4., 5.]])8.6 稀疏算子与存储类型推断
对稀疏数组有专门实现的算子集中在mx.nd.sparse下。例如dot(csr, dense):
shape = (3, 4) data = [7, 8, 9] indptr = [0, 2, 2, 3] indices = [0, 2, 1] a = mx.nd.sparse.csr_matrix((data, indices, indptr), shape=shape) rhs = mx.nd.ones((4, 1)) out = mx.nd.sparse.dot(a, rhs) # 调用专门针对 (csr, dense) 的稀疏 dot # [[ 15.], [ 0.], [ 9.]] <NDArray 3x1 @cpu(0)>存储类型推断规则:稀疏算子的输出存储类型由输入推断。例如a * 2的结果仍是 CSR(0 乘 2 还是 0),而a + ones((3,4))的结果变为稠密。可通过输出对象的.stype属性确认:
b = a * 2 c = a + mx.nd.ones(shape=(3, 4)) {'b.stype': b.stype, 'c.stype': c.stype} # {'b.stype': 'csr', 'c.stype': 'default'}存储类型回退(fallback):对未实现稀疏版本的稠密算子,仍可传入稀疏输入,但会有性能代价——MXNet 会把稀疏输入临时转成稠密再计算;若提供稀疏输出,则把稠密结果转回指定稀疏格式。回退发生时会打印警告信息(Jupyter 中显示在终端控制台):
e = mx.nd.sparse.zeros('csr', a.shape) d = mx.nd.log(a) # 稠密算子 + 稀疏输入 -> d.stype 为 'default' e = mx.nd.log(a, out=e) # 稠密算子 + 稀疏输出 -> e.stype 保持 'csr'8.7 稀疏数据加载
从 CSRNDArray 批量加载(mx.io.NDArrayIter):
data = mx.nd.array(np.arange(36).reshape((9, 4))).tostype('csr') labels = np.ones([9, 1]) batch_size = 3 dataiter = mx.io.NDArrayIter(data, labels, batch_size, last_batch_handle='discard') [batch.data[0] for batch in dataiter] # 每个 batch 都是 <CSRNDArray 3x4 @cpu(0)>从 libsvm 格式文件加载(mx.io.LibSVMIter)。libsvm 每行格式为<label> <col_idx1>:<value1> <col_idx2>:<value2> ...,例如矩阵有 6 列时,1 2:1.5 4:-3.5表示 label 为 1、数据为[[0, 0, 1.5, 0, -3.5, 0]]。注意列索引按行升序且从 0 开始(非 1 开始):
data_path = 'data.t' with open(data_path, 'w') as fout: fout.write('1.0 0:1 2:2\n') fout.write('1.0 0:3 5:4\n') fout.write('1.0 2:5 8:6 9:7\n') fout.write('1.0 3:8\n') fout.write('-1 0:0.5 9:1.5\n') fout.write('-2.0\n') fout.write('-3.0 0:-0.6 1:2.25 2:1.25\n') fout.write('-3.0 1:2 2:-1.25\n') fout.write('4 2:-1.2\n') data_train = mx.io.LibSVMIter(data_libsvm=data_path, data_shape=(10,), label_shape=(1,), batch_size=3) for batch in data_train: print(data_train.getdata()) # <CSRNDArray 3x10 @cpu(0)> print(data_train.getlabel()) # <NDArray 3 @cpu(0)>9. 稀疏 NDArray(二):RowSparseNDArray 与稀疏梯度更新
9.1 动机:稀疏梯度
训练大规模稀疏模型时,权重梯度往往也是稀疏的。设X为 1x2 矩阵、W为 2x3 矩阵,Y = dot(X, W):
import mxnet as mx X = mx.nd.array([[1, 0]]) W = mx.nd.array([[3, 4, 5], [6, 7, 8]]) Y = mx.nd.dot(X, W)则grad_W = dot(X.T, ones_like(Y))为:
grad_W = mx.nd.dot(X, mx.nd.ones_like(Y), transpose_a=True) # [[ 1. 1. 1.] # [ 0. 0. 0.]]由于X的第 1 列全为 0,grad_W的第 1 行也全为 0。真实世界中,与稀疏输入交互的参数其梯度通常存在大量全零行切片。稠密存储与计算这些 0 行是浪费;更重要的是,SGD、AdaGrad、Adam 等基于梯度的优化方法可以充分利用稀疏梯度实现高效更新。
RowSparseNDArray以行稀疏(row sparse)格式存储矩阵,专为“绝大多数行切片全为零”的数组设计,典型场景就是权重的稀疏梯度。
9.2 行稀疏格式
一个形状为[LARGE0, D1, ..., Dn]的多维 NDArray 用两个一维数组表示:
- data:任意 dtype、形状
[D0, D1, ..., Dn]; - indices:一维 int64 数组、形状
[D0],元素升序排列,存储非零行切片的行索引。
对应稠密数组满足:dense[rsp.indices[i], :, :, ...] = rsp.data[i, :, :, ...]。典型使用场景是LARGE0 >> D0且大多数行切片为零。
二维示例(5x3 矩阵,第 0、2 行非零):
data = [[1, 2, 3], [4, 0, 5]] indices = [0, 2]三维同样支持:一个 3x3x2 张量中第 0、1 个“行切片”非零,则data = [[[1,0],[0,2],[3,4]], [[5,0],[6,0],[0,0]]]、indices = [0, 1]。
RowSparseNDArray是NDArray的子类,其.stype属性值为"row_sparse"。
9.3 创建与检查
import mxnet as mx import numpy as np shape = (6, 2) data_list = [[1, 2], [3, 4]] indices_list = [1, 4] a = mx.nd.sparse.row_sparse_array((data_list, indices_list), shape=shape) # <RowSparseNDArray 6x2 @cpu(0)> b = mx.nd.sparse.row_sparse_array((np.array([[1, 2], [3, 4]]), np.array([1, 4])), shape=shape)row_sparse_array定义于 python/mxnet/ndarray/sparse.py#L1036。可用方法与 CSRNDArray 基本一致:.dtype、.asnumpy()、.data、.indices、.tostype、.cast_storage、.copy、.copyto。
a.asnumpy() # array([[ 0., 0.], # [ 1., 2.], # [ 0., 0.], # [ 0., 0.], # [ 3., 4.], # [ 0., 0.]], dtype=float32) {'a.stype': a.stype, 'data': a.data, 'indices': a.indices} # {'a.stype': 'row_sparse', 'data': [[1., 2.],[3., 4.]] <NDArray 2x2 @cpu(0)>, # 'indices': [1 4] <NDArray 2 @cpu(0)>}存储类型转换与 CSR 相同,tostype('row_sparse')/mx.nd.sparse.cast_storage(ones, 'row_sparse')均可实现default与row_sparse互转。
9.4 保留部分行切片:retain
mx.nd.sparse.retain(rsp, indices)按行索引保留指定行切片(其余置零):
data = [[1, 2], [3, 4], [5, 6]] indices = [0, 2, 3] rsp = mx.nd.sparse.row_sparse_array((data, indices), shape=(5, 2)) rsp_retained = mx.nd.sparse.retain(rsp, mx.nd.array([0, 1])) # 保留行 0、1:rsp_retained.asnumpy() 中仅第 0 行非零9.5 存储类型推断与回退
稀疏算子输出类型同样由输入推断。例如sparse.dot(csr, dense, transpose_a=True)的输出会被推断为row_sparse(因为转置点积产生行稀疏结构):
lhs = mx.nd.sparse.csr_matrix((data, indices, indptr), shape=(3, 5)) rhs = mx.nd.ones((3, 2)) transpose_dot = mx.nd.sparse.dot(lhs, rhs, transpose_a=True) # <RowSparseNDArray 5x2 @cpu(0)>标量运算保持行稀疏(a * 2结果为row_sparse),与稠密数组相加则退化为稠密(default)。对非稀疏算子,输入/输出会临时转成稠密再回退,并打印警告,行为与 CSRNDArray 一致。
9.6 稀疏优化器与 lazy update
当梯度为row_sparse存储且优化器以lazy_update=True创建时,MXNet 执行稀疏梯度更新。稀疏优化器只更新gradient.indices中出现的行切片对应的权重与状态。
以 SGD 为例,稠密更新规则为:
rescaled_grad = learning_rate * rescale_grad * clip(grad, clip_gradient) + weight_decay * weight state = momentum * state + rescaled_grad weight = weight - state而稀疏梯度的默认惰性更新(lazy update)为:
for row in grad.indices: rescaled_grad[row] = learning_rate * rescale_grad * clip(grad[row], clip_gradient) + weight_decay * weight[row] state[row] = momentum[row] * state[row] + rescaled_grad[row] weight[row] = weight[row] - state[row]注意:当weight_decay或momentum非零时,惰性更新与稠密更新的优化结果不同。如需禁用惰性更新,创建优化器时将lazy_update设为False。
实测示例:
shape = (4, 2) weight = mx.nd.ones(shape).tostype('row_sparse') data = [[1, 2], [4, 5]] indices = [1, 2] grad = mx.nd.sparse.row_sparse_array((data, indices), shape=shape) sgd = mx.optimizer.SGD(learning_rate=0.01, momentum=0.01) momentum = sgd.create_state(0, weight) sgd.update(0, weight, grad, momentum) # 只有行 1、2 的 weight 与 momentum 被更新: # weight = [[1, 1], [0.99, 0.98], [0.96, 0.95], [1, 1]] # momentum = [[0, 0], [-0.01, -0.02], [-0.04, -0.05], [0, 0]]目前 MXNet 中支持稀疏更新的优化器有mxnet.optimizer.SGD、mxnet.optimizer.Adam与mxnet.optimizer.AdaGrad。
9.7 GPU 支持
默认情况下,稀疏数组算子(包括 CSR 与 row_sparse)在 CPU 上执行。在 GPU 上创建需显式指定上下文,无 GPU 时把gpu_device改为mx.cpu():
import sys gpu_device = mx.gpu() # 无 GPU 时改为 mx.cpu() try: a = mx.nd.sparse.zeros('row_sparse', (100, 100), ctx=gpu_device) a except mx.MXNetError as err: sys.stderr.write(str(err))10. 学习路径与延伸阅读
本系列教程的完整目录(含各章节导航卡片)见 docs/python_docs/python/tutorials/packages/legacy/ndarray/index.rst,对应文档源文件均位于 docs/python_docs/python/tutorials/packages/legacy/ndarray/ 下:
- 01-ndarray-intro.md:NDArray 基础创建与属性;
- 02-ndarray-operations.md:运算、原地操作、切片、广播与 NumPy 互转;
- 03-ndarray-contexts.md:CPU/GPU 上下文管理;
- gotchas_numpy_in_mxnet.md:NumPy 使用中的常见误区与性能优化;
- sparse/csr.md 与 sparse/row_sparse.md:两种稀疏存储格式的完整教程。
如需深入源码,可重点阅读:
- python/mxnet/ndarray/ndarray.py:
NDArray类定义(#L266)及asnumpy(#L2635)、asscalar(#L2661)、copyto(#L2716)、as_in_context(#L2862)等核心方法; - python/mxnet/ndarray/sparse.py:稀疏工厂函数
csr_matrix(#L838)、row_sparse_array(#L1036)、zeros(#L1523)与array(#L1595); - src/engine/threaded_engine.cc:Execution Engine 的异步任务调度(
PushAsync,#L317),解释了 NDArray 运算为何异步非阻塞。
11. 结语
NDArray 是 MXNet 一切计算的基础:从最简单的数组创建、逐元素运算、原地写回,到 CPU/GPU 上下文迁移,再到为高维稀疏数据量身定制的 CSRNDArray 与 RowSparseNDArray,以及与之配合的惰性稀疏优化器。理解其异步执行模型与存储类型推断规则,是在 MXNet 上写出高性能训练代码的前提。如果你正在用新项目对接 MXNet,也可以留意官方后续推出的 NumPy 兼容数组类mx.np.ndarray,但经典 NDArray 所承载的存储抽象、上下文管理与异步执行思想依然一脉相承。
【免费下载链接】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),仅供参考