PRIMEAugmentation图像掩码增强报错:形状不匹配问题求助
解决PRIMEAugmentation增强时的形状不匹配RuntimeError
问题背景
实现图像与掩码的PRIMEAugmentation增强函数,已定义增强操作列表、PRIMEAugmentation类及调用逻辑,但运行时触发RuntimeError:
RuntimeError: output with shape [320, 320, 3] doesn't match the broadcast shape [320, 320, 320, 320, 3]
报错原因分析
- 输入维度理解错误:代码默认输入
x是带batch维度的格式,但实际传入的是单张图片(形状为[H,W,C],比如[320,320,3]),导致x.shape[0]取到的是图像高度而非batch size,后续Dirichlet采样、操作索引生成全部错位。 - 增强操作调用错误:
self.augmentations[op]中op是长度为x.shape[0]的索引列表,直接用列表索引取增强函数会触发维度错误,应为每个样本单独选择操作并执行。 - 张量维度不匹配:初始化的
x_aug是单张图形状,但后续计算时给x_i添加了多余维度(unsqueeze(1)),同时权重的维度扩展错误,导致广播时形状冲突。 - 类未绑定增强操作列表:PRIMEAugmentation类的
__init__方法未接收并保存外部定义的augmentations列表,self.augmentations未定义会触发后续调用错误。
修复后的代码
1. 修正PRIMEAugmentation类
import torch from torch.distributions import Dirichlet, Beta class PRIMEAugmentation: def __init__(self, augmentations, mixture_width=3, mixture_depth=-1): self.augmentations = augmentations # 绑定外部传入的增强操作列表 self.mixture_width = mixture_width # 默认mixture_depth等于mixture_width,避免负数值导致循环异常 self.mixture_depth = mixture_depth if mixture_depth != -1 else mixture_width def __call__(self, x, mask): # 处理单张输入:自动添加batch维度,转为[1, H, W, C]格式 if len(x.shape) == 3: x = torch.from_numpy(x).unsqueeze(0) mask = torch.from_numpy(mask).unsqueeze(0) else: x = torch.from_numpy(x) mask = torch.from_numpy(mask) batch_size = x.shape[0] # 采样Dirichlet权重:形状为[batch_size, mixture_width] ws = Dirichlet(torch.ones(self.mixture_width)).sample((batch_size,)) # 采样Beta混合系数:扩展为[batch_size, 1, 1, 1],适配后续广播 m = Beta(torch.ones(1), torch.ones(1)).sample().expand(batch_size, 1, 1, 1) 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() for _ in range(self.mixture_depth): # 为每个样本随机选择一个增强操作并执行 for idx in range(batch_size): op_idx = torch.randint(len(self.augmentations), size=(1,)).item() x_i[idx], mask_i[idx] = self.augmentations[op_idx](x_i[idx], mask_i[idx]) # 调整权重维度,与x_i广播匹配:[batch_size,1,1,1] weight = ws[:, i].unsqueeze(-1).unsqueeze(-1).unsqueeze(-1) x_aug += weight * x_i mask_aug += weight * mask_i # 混合原始图与增强后的图 mixed = (1 - m) * x + m * x_aug mixed_mask = (1 - m) * mask + m * mask_aug # 若输入是单张图,移除batch维度后转回numpy if batch_size == 1: return mixed.squeeze(0).numpy(), mixed_mask.squeeze(0).numpy() return mixed.numpy(), mixed_mask.numpy()
2. 修正调用逻辑
# 初始化增强器时传入预定义的增强操作列表 augmenter_PRIMEAugmentation = aug_lib_new.PRIMEAugmentation(augmentations) 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: # 直接调用增强器实例,避免eval的不安全写法 if aug_trans == "PRIMEAugmentation": image, mask = augmenter_PRIMEAugmentation(image, mask)
关键修复点说明
- 新增单图输入兼容:自动添加/移除batch维度,同时支持单图和批量输入场景。
- 修正权重采样维度:基于实际batch size生成Dirichlet权重,解决维度错位问题。
- 逐个样本执行增强:避免批量索引操作的错误,确保每个样本随机应用增强策略。
- 调整张量维度匹配:优化权重的维度扩展方式,避免多余维度导致的广播冲突。
- 移除不安全的
eval调用:通过条件判断直接调用增强器实例,提升代码安全性与可读性。
内容的提问来源于stack exchange,提问作者sccomp
相关产品推荐
相关产品推荐

