PyTorch中KMNIST数据集张量灰度通道提取与维度调整问题
KMNIST张量维度调整与灰度通道查看方案
1. 明确张量各维度含义
你当前的张量形状torch.Size([60000, 28, 28])中,60000是样本数(即你说的row_number),后两个28分别是单张灰度图的高度和宽度。灰度图默认通道数为1,但这个维度在原始数据中被省略了,所以没显式体现出来。
2. 调整维度到(row_number, height, width, channel)
不需要移动现有轴,直接给张量添加灰度通道维度即可,用PyTorch的unsqueeze方法在最后一维插入通道维度:
# 从(N, H, W)格式转为(N, H, W, C),C=1(灰度通道) x_train_with_channel = x_train.unsqueeze(-1) # 确认新形状 print(x_train_with_channel.shape) # 输出: torch.Size([60000, 28, 28, 1])
你之前用np.moveaxis(x_train,0, -1)是错误的,因为把样本轴移到了最后,完全偏离了需求。
3. 查看通道及其数值
- 查看单张样本的通道数据:
取指定样本的通道张量,再通过索引查看具体像素的通道数值:# 取第2个样本(索引为1)的通道数据,形状为(28, 28, 1) sample_channel = x_train_with_channel[1] # 查看(10,10)位置的灰度通道数值 print(sample_channel[10, 10, 0]) - 查看整个通道的数值矩阵:
可以转成NumPy数组后直接打印,更直观:sample_channel_np = sample_channel.numpy() print(sample_channel_np) # 输出28x28x1的灰度值矩阵 - 验证通道数:你用
torchvision.transforms.functional.get_image_num_channels(x_train[1])得到1是正确的,因为原始灰度图确实只有1个通道,只是原张量没显式包含这个维度。
内容的提问来源于stack exchange,提问作者learner1
相关产品推荐
相关产品推荐

