如何用Python NumPy读取多幅3D图像并调整数组维度适配模型输入
维度不匹配的解决方法
你当前拼接得到的数组维度顺序为(batch_size, 模态通道数, 深度, 高度, 宽度),模型要求的输入维度顺序为(batch_size, 深度, 高度, 宽度, 模态通道数),本质是通道维度的位置不对,同时原有拼接逻辑存在顺序反转、效率低的问题,按下面方法调整即可。
方法1:直接修改加载逻辑(推荐)
不要用前置np.append的方式拼接数组,改用列表暂存每个模态处理后的影像,最后直接在最后一个轴堆叠,单样本输出形状直接为(365, 256, 256, 3),经DataLoader组batch后自然得到(batch_size, 365, 256, 256, 3)的符合要求的形状,同时不会出现模态顺序反转的问题。
修正后的完整代码如下:
def Load_function(path): f_img = nib.load(path) img_data = f_img.get_fdata() return img_data def __load__(self, id_name): image_path = os.path.join(self.path, id_name) image_list = [] ## 逐张读取3个模态的影像 for imname in ["image2B.nii.gz", "image1to2_nlB.nii.gz", "diffFSL.nii.gz"]: img = Load_function(os.path.join(image_path, imname)) img = resize_data(img) ## 归一化 img = img / np.percentile(img, 99.5) image_list.append(img) ## 在最后一个维度拼接3个模态通道 image = np.stack(image_list, axis=-1) ## 读取掩码 mask = Load_function(os.path.join(image_path, "ground_truth.nii.gz")) mask = resize_data(mask) return image, mask
方法2:对已生成的错误形状数组做维度转换
如果你不想修改现有加载逻辑,已经拿到了形状为(2, 3, 365, 256, 256)的批量数组,只需要加一行代码移动维度位置即可:
# 将位置为1的通道维度,移动到数组最后一位 image = np.moveaxis(image, source=1, destination=-1)
转换后数组形状直接变为(2, 365, 256, 256, 3),完全匹配模型输入要求。
原有代码问题说明
- 每次执行
np.append(img[np.newaxis, ...],image, axis=0)会把新读取的影像插到数组最前端,最终3个模态的存储顺序和读取顺序相反 - 拼接时把通道维度放在了batch维之后、空间维之前,不符合模型通道维居尾的格式要求
- 循环中反复用
np.append拼接数组会反复开辟新内存,加载效率很低
内容的提问来源于stack exchange,提问作者Ehsan Alavi
相关产品推荐
相关产品推荐

