U-Net可视化预测掩码时维度错误:TypeError(无效形状(2023,2023,256))
U-Net预测结果可视化报错解决
我已经完成U-Net网络训练,现在尝试可视化预测结果。我认为掩码应该和原图尺寸一致且为单通道,这个理解对吗?以下是相关代码和报错信息:
加载模型
# 加载模型权重 weights_path = unet_dir + "unet1.pt" device = "cpu" unet = UNet(in_channels=3, out_channels=3, init_features=8) unet.to(device) unet.load_state_dict(torch.load(weights_path, map_location=device))
初始化函数
# 定义数据增强 inference_transform = A.Compose([ A.Resize(256, 256, always_apply=True), A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)), ToTensorV2() ]) # 定义预测函数 def predict(model, img, device): model.eval() with torch.no_grad(): images = img.to(device) output = model(images) predicted_masks = (output.squeeze() >= 0.5).float().cpu().numpy() return(predicted_masks) # 定义加载图片并输出掩码的函数 def get_mask(img_path): image = cv2.imread(img_path) # assert image is not None image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) original_height, original_width = tuple(image.shape[:2]) image_trans = inference_transform(image = image) image_trans = image_trans["image"] image_trans = image_trans.unsqueeze(0) image_mask = predict(unet, image_trans, device) # image_mask = image_mask.astype(np.int16) image_mask = cv2.resize(image_mask,(original_width, original_height), interpolation=cv2.INTER_NEAREST) # image_mask = cv2.resize(image_mask, (original_height, original_width)) # Y_train[n] = mask > 0.5 return(image_mask)
测试代码
# 测试图片路径 example_path = "../input/test-image/10078.tiff" image = cv2.imread(example_path) # assert image is not None image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask = get_mask(example_path) # masked_img = image*np.expand_dims(mask, 2).astype("uint8") # 绘制原图、掩码及叠加图 fig, (ax1, ax2) = plt.subplots(2) ax1.imshow(image) ax2.imshow(mask) # ax3.imshow(masked_img)
报错信息
--------------------------------------------------------------------------- TypeError Traceback (most recent call last) /tmp/ipykernel_4859/3003834023.py in <module> 13 14 ax1.imshow(image) ---> 15 ax2.imshow(mask) 16 #ax3.imshow(masked_img) /opt/conda/lib/python3.7/site-packages/matplotlib/_api/deprecation.py in wrapper(*args, **kwargs) 457 "parameter will become keyword-only %(removal)s.", 458 name=name, obj_type=f"parameter of {func.__name__}()") ---> 459 return func(*args, **kwargs) 460 461 # Don't modify *func*'s signature, as boilerplate.py needs it. /opt/conda/lib/python3.7/site-packages/matplotlib/__init__.py in inner(ax, data, *args, **kwargs) 1412 def inner(ax, *args, data=None, **kwargs): 1413 if data is None: -> 1414 return func(ax, *map(sanitize_sequence, args), **kwargs) 1415 1416 bound = new_sig.bind(ax, *args, **kwargs) /opt/conda/lib/python3.7/site-packages/matplotlib/axes/_axes.py in imshow(self, X, cmap, norm, aspect, interpolation, alpha, vmin, vmax, origin, extent, interpolation_stage, filternorm, filterrad, resample, url, **kwargs) 5485 **kwargs) 5486 -> 5487 im.set_data(X) 5488 im.set_alpha(alpha) 5489 if im.get_clip_path() is None: /opt/conda/lib/python3.7/site-packages/matplotlib/image.py in set_data(self, A) 714 or self._A.ndim == 3 and self._A.shape[-1] in [3, 4]): 715 raise TypeError("Invalid shape {} for image data" -> 716 .format(self._A.shape)) 717 718 if self._A.ndim == 3: TypeError: Invalid shape (2023, 2023, 256) for image data
问题分析与解决方法
核心问题
报错显示掩码形状为(2023, 2023, 256),根源在于两点:一是U-Net输出通道数设为3,不符合单通道掩码的需求;二是cv2.resize处理多通道数据时改变了维度顺序,导致最终形状不满足matplotlib.imshow的要求(单通道需为(H,W),3通道需为(H,W,3))。
解决步骤
调整模型输出通道数:
若为二分类任务,U-Net的out_channels应设为1,修改模型初始化代码:unet = UNet(in_channels=3, out_channels=1, init_features=8)若为多分类任务,需根据类别数设置
out_channels,后续再将多通道输出转为单通道掩码。修正掩码处理逻辑:
即使暂时不修改模型,也需要将多通道输出转为单通道。示例如下:# 在predict函数中修改 def predict(model, img, device): model.eval() with torch.no_grad(): images = img.to(device) output = model(images) # 二分类取单通道,多分类取概率最大的类别 if output.shape[1] == 1: predicted_masks = (output.squeeze() >= 0.5).float().cpu().numpy() else: predicted_masks = torch.argmax(output, dim=1).squeeze().cpu().numpy() return(predicted_masks)确认掩码维度:
确保最终返回的掩码是二维数组(H,W),matplotlib.imshow才能正常显示。多分类单通道掩码可通过指定cmap参数区分类别,比如ax2.imshow(mask, cmap='viridis')。
关于掩码的理解
你的理解是正确的:掩码应与原图尺寸一致,二分类任务下为单通道(每个像素表示是否为目标);多分类任务下可以是单通道(每个像素为类别索引)或多通道(每个通道对应一个类别的概率)。
内容的提问来源于stack exchange,提问作者Alain Michael Janith Schroter
相关产品推荐
相关产品推荐

