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

PRIME数据增强后触发AttributeError:Image对象无shape属性

PRIME数据增强后Albumentations Normalize报错解决

问题概述

实现PRIME数据增强时,先因输入为4维张量触发ValueError: pic should be 2/3 dimensional. Got 4 dimensions.,处理后(移除批量维度、转(H,W,C)格式numpy数组、转PIL Image),又出现AttributeError: 'Image' object has no attribute 'shape',报错指向Albumentations的Normalize操作。

核心原因

Albumentations所有变换仅支持numpy数组格式的图像输入,而你返回的是PIL Image对象——PIL Image没有shape属性,导致Normalize在获取图像行列数时失败。

解决方案

直接返回(H,W,C)格式的numpy数组,无需转为PIL Image。若后续流程需要PIL Image,应在Albumentations所有变换完成后再转换。

修改后的PrimeAugment代码

import numpy as np
import torchvision.transforms as transforms

class PrimeAugment:
    def __init__(self, prime_module):
        self.prime_module = prime_module

    def __call__(self, img, mask):
        # 转张量并添加批量维度供PRIME处理
        img = transforms.ToTensor()(img).unsqueeze(0)
        img_prime = self.prime_module(img)
        print("img_prime shape:", img_prime.shape)
        
        # 移除批量维度,转numpy数组并调整为(H,W,C)格式
        img_prime = img_prime[0, :, :, :]
        img_prime = img_prime.detach().cpu().numpy().transpose(1, 2, 0)
        img_prime = (img_prime * 255).astype(np.uint8)
        print("processed image shape:", img_prime.shape)
        
        # 直接返回numpy数组,适配Albumentations
        return img_prime, mask

额外适配提示

若输入到PrimeAugment的是numpy数组而非PIL Image,可将开头的张量转换代码替换为:

# 输入为(H,W,C)格式numpy数组时的转换方式
img = torch.from_numpy(img.transpose(2, 0, 1)).unsqueeze(0).float() / 255.0

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 16:57:41