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

实现PRIMEAugmentation图像掩码增强时遇维度错误求助

PRIMEAugmentation维度错误排查:ValueError: pic should be 2/3 dimensional. Got 4 dimensions.

问题背景

正在为图像分割任务实现PRIMEAugmentation增强函数,相关代码如下:

增强操作定义

if config.AUG == "PRIMEAugmentation":
    augmentations = [autocontrast, equalize, posterize, rotate, solarize, shear_x, shear_y, translate_x, translate_y]

PRIME增强类实现

import torch
from torch.distributions import Dirichlet, Beta

class PRIMEAugmentation:
    def __init__(self, mixture_width=3, mixture_depth=-1):
        self.mixture_width = mixture_width
        self.mixture_depth = mixture_depth

    def __call__(self, x, mask):
        x = torch.from_numpy(x).to(torch.float32)
        mask = torch.from_numpy(mask)
        ws = Dirichlet(torch.ones(self.mixture_width)).sample((x.shape[0],))
        m = Beta(torch.ones(1), torch.ones(1)).sample().expand(x.shape[0], 1, 1, 1)

        x_aug = torch.zeros_like(x).to(torch.float32)
        mask_aug = torch.zeros_like(mask).to(torch.float32)
        for i in range(self.mixture_width):
            x_i = x.clone()
            mask_i = mask.clone()
            for d in range(self.mixture_depth):
                op = torch.randint(len(self.augmentations), size=(x.shape[0],)).tolist()
                x_i, mask_i = self.augmentations[op](x_i, mask_i)
            print("ws[:, i] shape:", ws[:, i].shape)
            print("x_i shape:", x_i.shape)
            print("mask_i shape:", mask_i.shape)
            x_aug += ws[:, i][:, None, None] * x_i.to(torch.float32)
            mask_aug += ws[:, i][:, None] * mask_i.to(torch.float32)

        mixed = (1 - m) * x + m * x_aug.sum(dim=1)
        mixed_mask = (1 - m) * mask + m * mask_aug.sum(dim=1)
        return mixed.numpy().astype(np.uint8), mixed_mask.numpy().astype(np.uint8)

调用方式

augmenter_PRIMEAugmentation = aug_lib_new.PRIMEAugmentation()

import os

def image_mask_transformation(image,mask,img_trans,aug_trans=False):
    transformed = img_trans(image=image, mask=mask)
    image = transformed["image"]
    mask = transformed["mask"]

    if aug_trans in augmenter_list:
        image,mask = eval('augmenter_'+aug_trans)(image, mask)

数据集调用类

class SegmentationDataset(Dataset):
    def __init__(self, imagePaths, maskPaths, img_trans, aug_trans = False, baug = 1):
        self.imagePaths = imagePaths
        self.maskPaths = maskPaths
        self.img_trans = img_trans
        self.aug_trans = aug_trans
        self.baug = baug

    def __len__(self):
        return len(self.imagePaths)

    def __getitem__(self, idx):
        imagePath = self.imagePaths[idx]

        image = cv2.imread(imagePath)
        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
        mask = cv2.imread(self.maskPaths[idx], 0)

        image_store, mask_store = image_mask_transformation(image, mask, self.img_trans, self.aug_trans)
        return image_store, mask_store

错误详情

运行时抛出维度不匹配错误:

raise ValueError('pic should be 2/3 dimensional. Got {} dimensions.'.format(pic.ndim))
ValueError: pic should be 2/3 dimensional. Got 4 dimensions.

完整错误栈:

Traceback (most recent call last):
  File "/home/Crack-PRIME4/main.py", line 422, in <module>
    train_logs = train_step(model, optimizer, criteria, trainLoader, accumulation_steps, scaler, epoch, epochs)
  File "/home/Crack-PRIME4/main.py", line 240, in train_step
    for idx, data in enumerate(bar):
  File "/home/anaconda3/envs/myenv/lib/python3.9/site-packages/tqdm/std.py", line 1182, in __iter__
    for obj in iterable:
  File "/home//anaconda3/envs/myenv/lib/python3.9/site-packages/torch/utils/data/dataloader.py", line 630, in __next__
    data = self._next_data()
  File "/home/anaconda3/envs/myenv/lib/python3.9/site-packages/torch/utils/data/dataloader.py", line 1345, in _next_data
    return self._process_data(data)
  File "/home/anaconda3/envs/myenv/lib/python3.9/site-packages/torch/utils/data/dataloader.py", line 1371, in _process_data
    data.reraise()
  File "/home/anaconda3/envs/myenv/lib/python3.9/site-packages/torch/_utils.py", line 694, in reraise
    raise exception
ValueError: Caught ValueError in DataLoader worker process 0.
Original Traceback (most recent call last):
  File "/home/anaconda3/envs/myenv/lib/python3.9/site-packages/torch/utils/data/_utils/worker.py", line 308, in _worker_loop
    data = fetcher.fetch(index)
  File "/home/anaconda3/envs/myenv/lib/python3.9/site-packages/torch/utils/data/_utils/fetch.py", line 51, in fetch
    data = [self.dataset[idx] for idx in possibly_batched_index]
  File "/home/anaconda3/envs/myenv/lib/python3.9/site-packages/torch/utils/data/_utils/fetch.py", line 51, in <listcomp>
    data = [self.dataset[idx] for idx in possibly_batched_index]
  File "/home/Crack-PRIME4/tool/dataset.py", line 215, in __getitem__
    image_store, mask_store = image_mask_transformation(image, mask, self.img_trans, self.aug_trans)
  File "/home/Crack-PRIME4/tool/dataset.py", line 188, in image_mask_transformation
    final_image = transforms.ToTensor()(image)
  File "/home/anaconda3/envs/myenv/lib/python3.9/site-packages/torchvision/transforms/transforms.py", line 97, in __call__
    return F.to_tensor(pic)
  File "/home/anaconda3/envs/myenv/lib/python3.9/site-packages/torchvision/transforms/functional.py", line 105, in to_tensor
    raise ValueError('pic should be 2/3 dimensional. Got {} dimensions.'.format(pic.ndim))
ValueError: pic should be 2/3 dimensional. Got 4 dimensions.

原因分析

  1. 批量与单样本维度不匹配:PRIME增强类的__call__方法设计为处理批量数据(4D张量:(batch_size, H, W, C)),但数据集__getitem__传入的是单张图像(3D:(H, W, C))和单张掩码(2D:(H, W))。此时x.shape[0]是图像高度而非batch_size,导致后续权重生成和张量操作错误地增加了一个维度,最终输出4D数据。
  2. 实例属性缺失:增强类__init__未将全局定义的augmentations列表赋值给self.augmentations,运行时会触发AttributeError,属于潜在问题。
  3. 增强操作调用错误:self.augmentations[op]中op是批量索引列表,单样本场景下无法正确索引单个增强操作。

解决方法

1. 修正PRIME增强类,适配单样本输入

修改类逻辑,处理无batch维度的单样本数据:

import torch
import numpy as np
from torch.distributions import Dirichlet, Beta

class PRIMEAugmentation:
    def __init__(self, mixture_width=3, mixture_depth=-1, augmentations=None):
        self.mixture_width = mixture_width
        self.mixture_depth = mixture_depth
        self.augmentations = augmentations  # 接收外部传入的增强操作列表

    def __call__(self, x, mask):
        # x: (H, W, C) 单张RGB图像;mask: (H, W) 单张灰度掩码
        x = torch.from_numpy(x).to(torch.float32)
        mask = torch.from_numpy(mask).to(torch.float32)
        
        # 生成单样本的混合权重
        ws = Dirichlet(torch.ones(self.mixture_width)).sample()
        m = Beta(torch.ones(1), torch.ones(1)).sample().item()

        x_aug = torch.zeros_like(x)
        mask_aug = torch.zeros_like(mask)
        
        for i in range(self.mixture_width):
            x_i = x.clone()
            mask_i = mask.clone()
            # mixture_depth为-1时,随机选择1-3层增强(参考PRIME原文)
            depth = np.random.randint(1, 4) if self.mixture_depth == -1 else self.mixture_depth
            for _ in range(depth):
                # 随机选择单个增强操作
                op_idx = torch.randint(len(self.augmentations), size=(1,)).item()
                x_i, mask_i = self.augmentations[op_idx](x_i, mask_i)
            # 加权累加增强结果
            x_aug += ws[i] * x_i
            mask_aug += ws[i] * mask_i

        # 混合原始与增强数据,截断到0-255范围
        mixed = (1 - m) * x + m * x_aug
        mixed_mask = (1 - m) * mask + m * mask_aug
        
        return np.clip(mixed.numpy(), 0, 255).astype(np.uint8), np.clip(mixed_mask.numpy(), 0, 255).astype(np.uint8)

2. 修正增强器初始化与调用

初始化时传入增强操作列表,避免全局依赖:

# 定义增强操作列表
augmentations = [autocontrast, equalize, posterize, rotate, solarize, shear_x, shear_y, translate_x, translate_y]
# 初始化增强器并传入列表
augmenter_PRIMEAugmentation = aug_lib_new.PRIMEAugmentation(augmentations=augmentations)

3. 增加维度校验(可选)

在image_mask_transformation函数中增加维度断言,提前发现问题:

def image_mask_transformation(image,mask,img_trans,aug_trans=False):
    transformed = img_trans(image=image, mask=mask)
    image = transformed["image"]
    mask = transformed["mask"]

    if aug_trans in augmenter_list:
        image,mask = eval('augmenter_'+aug_trans)(image, mask)
    # 确保图像和掩码维度符合要求
    assert len(image.shape) == 3, f"图像应为3D,当前为{len(image.shape)}D"
    assert len(mask.shape) == 2, f"掩码应为2D,当前为{len(mask.shape)}D"
    return image, mask

关键注意点

  • PRIME原始实现针对批量数据,需根据使用场景(单样本/批量)调整维度逻辑。
  • 增强操作函数需支持处理单样本张量,并同时返回增强后的图像和掩码。
  • 转换回numpy数组时必须用np.clip限制值范围,避免溢出导致异常。

内容的提问来源于stack exchange,提问作者sccomp

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 12:35:54