PyTorch自定义4D图像Dataset传入Conv3D出现类型不匹配报错如何解决
错误原因
这个报错的核心是数据类型不匹配:
- 你加载PNG图像时指定了
.astype(np.uint8),转成PyTorch张量后就是torch.ByteTensor类型,值域为0~255 - PyTorch所有卷积层的权重、偏置默认都是
torch.FloatTensor(float32)类型,要求输入数据也必须是浮点类型,因此触发了类型不匹配报错。
修复后的自定义Dataset代码
除了修正数据类型,还补充了常规的图像归一化操作,避免训练时数值范围过大导致梯度异常,同时修正了原代码中可能导致维度异常的嵌套数组写法,补全了缺失的依赖导入:
import torch import os import numpy as np from PIL import Image import pandas as pd from torch.utils.data import Dataset class TrainImages(Dataset): def __init__(self, csv_file, root_dir, transform=None): self.annotations = pd.read_csv(csv_file) self.root_dir = root_dir self.transform = transform def __len__(self): return len(self.annotations) def __getitem__(self, index): # 拼接当前扫描对应的文件夹路径 img_folder = os.path.join(self.root_dir, str(self.annotations.iloc[index, 0]).zfill(5)) # 逐个加载31张RGB图像 imgs_list = [] for i in range(31): img_path = os.path.join(img_folder, f"rgb-{i}.png") imgs_list.append(np.array(Image.open(img_path))) # 堆叠+调整维度为 (Channels, Depth, Height, Width) 符合Conv3D输入要求 imgs_arr = np.stack(imgs_list, axis=0).transpose(3,0,1,2) # 转float32+归一化到0~1区间,是修复报错的核心步骤 imgs_arr = imgs_arr.astype(np.float32) / 255.0 # 转PyTorch张量 image = torch.from_numpy(imgs_arr) # 执行自定义增广(如果有传入) if self.transform: image = self.transform(image) y_label = torch.tensor(int(self.annotations.iloc[index, 1])) return (image, y_label)
修复验证
改完后运行你的测试代码即可正常执行:
import torch.nn as nn x = torch.unsqueeze(dataset[0][0], 0) print(x.shape, x.dtype) # 输出应为 torch.Size([1, 3, 31, 512, 512]) torch.float32 m = nn.Conv3d(3,12,3) out = m(x) print(out.shape) # 正常输出卷积后的张量尺寸即代表修复完成
内容的提问来源于stack exchange,提问作者durian
相关产品推荐
相关产品推荐

