实现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.
原因分析
- 批量与单样本维度不匹配:PRIME增强类的
__call__方法设计为处理批量数据(4D张量:(batch_size, H, W, C)),但数据集__getitem__传入的是单张图像(3D:(H, W, C))和单张掩码(2D:(H, W))。此时x.shape[0]是图像高度而非batch_size,导致后续权重生成和张量操作错误地增加了一个维度,最终输出4D数据。 - 实例属性缺失:增强类
__init__未将全局定义的augmentations列表赋值给self.augmentations,运行时会触发AttributeError,属于潜在问题。 - 增强操作调用错误:
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
相关产品推荐
相关产品推荐

