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

PRIMEAugmentation图像掩码增强报错:形状不匹配问题求助

解决PRIMEAugmentation增强时的形状不匹配RuntimeError

问题背景

实现图像与掩码的PRIMEAugmentation增强函数,已定义增强操作列表、PRIMEAugmentation类及调用逻辑,但运行时触发RuntimeError:

RuntimeError: output with shape [320, 320, 3] doesn't match the broadcast shape [320, 320, 320, 320, 3]

报错原因分析

  1. 输入维度理解错误:代码默认输入x是带batch维度的格式,但实际传入的是单张图片(形状为[H,W,C],比如[320,320,3]),导致x.shape[0]取到的是图像高度而非batch size,后续Dirichlet采样、操作索引生成全部错位。
  2. 增强操作调用错误:self.augmentations[op]中op是长度为x.shape[0]的索引列表,直接用列表索引取增强函数会触发维度错误,应为每个样本单独选择操作并执行。
  3. 张量维度不匹配:初始化的x_aug是单张图形状,但后续计算时给x_i添加了多余维度(unsqueeze(1)),同时权重的维度扩展错误,导致广播时形状冲突。
  4. 类未绑定增强操作列表: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 12:04:58