PyTorch绘制MNIST图像时的Tensor索引相关疑问
PyTorch MNIST绘图张量索引问题解答
维度结构说明
你从DataLoader中取出的批量数据example_data是四维张量,维度顺序为[批量大小, 通道数, 图像高度, 图像宽度],对应你的参数设置,实际shape为[100, 1, 28, 28],各维度的索引从0开始计数。
索引写法的含义解释
example_data[j]是对四维张量的**第一个维度(批量样本维度)**做索引,取出批量中第j个样本,得到的是一个三维张量,shape为[1, 28, 28],三个维度分别对应单样本的「通道、高度、宽度」。example_data[j][0]是对刚才取出的三维单样本张量再做索引,取它**第一个维度(通道维度)**的第0个元素,最终得到shape为[28, 28]的二维数组,正好匹配plt.imshow的输入要求,这里的0确实是对通道维度做索引。
为什么不能写成example_data[0,j]
逗号分隔的索引是直接作用于最外层的原始张量的,example_data[0,j]的含义是:
- 第一个索引0取批量维度的第0个样本
- 第二个索引j取通道维度的第j个元素
你的通道维度只有1个元素,最大索引为0,j的取值范围是0~99,这么写要么直接报索引越界,要么拿到的完全不是你想要的第j个样本的图像数据。
如果你想用更高效的逗号分隔写法(避免生成中间张量),正确的写法是example_data[j, 0],它和example_data[j][0]的运行结果完全一致。
补充说明
你用到的np.squeeze(example_data[j])本质是自动删除所有长度为1的维度,刚好这个场景里通道维度长度是1,所以也能得到[28, 28]的输出,和取第0个通道的效果一致。
内容的提问来源于stack exchange,提问作者Nordcus
相关产品推荐
相关产品推荐

