You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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]的含义是:

  1. 第一个索引0取批量维度的第0个样本
  2. 第二个索引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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.27 04:15:09