如何在MXNet中提取典型行?及基于指定索引选取批量数据
在MXNet中提取指定行/元素的方法
一、通用提取典型行的方式
在MXNet里,提取数组的指定行(或者高维数组中某一轴的指定元素),常用这几种方法:
- 直接索引:如果是二维数组,直接用下标索引就能搞定,比如想提取第0和第2行,就写
data[[0, 2]];要是高维数组,也可以用类似的多维索引,比如data[:, [0,1], :](对应你例子里的需求)。 nd.slice_axis函数:适合提取连续的行/元素,比如想提取第1到第3行(左闭右开),可以写mx.nd.slice_axis(data, axis=0, begin=1, end=3),axis参数指定要操作的轴。nd.take和nd.pick函数:这俩是提取非连续、自定义索引的利器,尤其适合高维批量数据的场景,也是解决你问题的关键。
二、针对你的批量数据提取需求
先看你给出的代码:
import mxnet as mx data = mx.nd.array(range(24)).reshape(2,3,4) index = mx.nd.array([[0,1],[1,2]])
先明确下数据结构:data是形状为(2,3,4)的批量数据,也就是2个样本,每个样本有3行,每行4个特征;index是每个样本要提取的行索引,样本0取第0、1行,样本1取第1、2行。
用nd.take实现
take函数需要指定要操作的轴(这里是行所在的轴axis=1),然后传入索引数组即可:
result_take = mx.nd.take(data, index, axis=1) print(result_take)
运行后得到的结果形状是(2,2,4),正好是每个样本提取2行,每行4个特征,内容如下:
[[[ 0. 1. 2. 3.] [ 4. 5. 6. 7.]] [[16. 17. 18. 19.] [20. 21. 22. 23.]]]
用nd.pick实现
pick其实是take的简化版,当你的索引数组维度和原数组除了目标轴之外的维度完全匹配时,用pick更直观:
result_pick = mx.nd.pick(data, index, axis=1) print(result_pick)
得到的结果和take完全一致,因为pick会自动对齐批量维度的索引。
小提示
如果你的索引是一维的,比如想给所有样本都提取第0、2行,那直接写index = mx.nd.array([0,2]),然后用take或pick指定axis=1就行,结果会是(2,2,4)的形状。
内容的提问来源于stack exchange,提问作者partida
相关产品推荐
相关产品推荐

