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

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))。

解决步骤

  1. 调整模型输出通道数:
    若为二分类任务,U-Net的out_channels应设为1,修改模型初始化代码:

    unet = UNet(in_channels=3, out_channels=1, init_features=8)
    

    若为多分类任务,需根据类别数设置out_channels,后续再将多通道输出转为单通道掩码。

  2. 修正掩码处理逻辑:
    即使暂时不修改模型,也需要将多通道输出转为单通道。示例如下:

    # 在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)
    
  3. 确认掩码维度:
    确保最终返回的掩码是二维数组(H,W),matplotlib.imshow才能正常显示。多分类单通道掩码可通过指定cmap参数区分类别,比如ax2.imshow(mask, cmap='viridis')。

关于掩码的理解

你的理解是正确的:掩码应与原图尺寸一致,二分类任务下为单通道(每个像素表示是否为目标);多分类任务下可以是单通道(每个像素为类别索引)或多通道(每个通道对应一个类别的概率)。

内容的提问来源于stack exchange,提问作者Alain Michael Janith Schroter

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 11:36:14