调用torchvision.transforms.ColorJitter报通道错误的解决方法
报错根因
触发该通道数报错的直接原因是输入Tensor的维度顺序不符合torchvision变换接口的要求,同时代码还存在数据格式不匹配、保存逻辑错误的问题,具体如下:
torchvision.transforms系列算子接收Tensor格式图像时,强制要求维度顺序为(通道数C, 图像高度H, 图像宽度W),也就是CHW格式。你传入的Tensor形状为(1080, 1920, 3),是numpy/OpenCV默认的HWC(高、宽、通道)顺序,算子误将第一维的高度值1080识别为通道数,才会抛出"找到1080个通道,仅支持1/3通道"的错误。ColorJitter处理Tensor输入时,要求输入为torch.float32类型、像素值范围在[0.0, 1.0]区间,直接从numpy数组转来的Tensor默认是uint8类型、像素值范围[0,255],即使调整对维度顺序,也会出现色彩计算异常的问题。- 保存逻辑存在两处错误:一是
cv2.imwrite不支持直接接收PyTorch Tensor、也不支持直接传入存储4张图像的列表作为输入;二是OpenCV接口默认使用BGR通道顺序,直接传入torchvision输出的RGB格式图像会出现蓝红色彩反转。
修复方案
按以下顺序调整代码即可正常运行:
- 维度转换:将输入Tensor从HWC顺序调整为CHW顺序,可使用
permute(2,0,1)方法完成维度换位。 - 数值归一化:将Tensor转为
float32类型,将像素值从[0,255]线性缩放至[0.0,1.0]区间;如果原始图像是OpenCV读取的,需要先把BGR通道顺序转为RGB顺序再输入变换算子。 - 结果后处理:变换完成后,将每张输出Tensor的维度从CHW转回HWC格式,转为numpy数组,把像素值缩放回
[0,255]区间并转为uint8类型,再把RGB通道顺序转回OpenCV支持的BGR顺序,逐张调用保存接口即可。
修正后可直接运行的代码
import torch import torchvision import cv2 as cv import numpy as np # 此处image为你原有的形状(1080,1920,3)的输入数组,若为OpenCV读取默认是BGR顺序,先转RGB image_rgb = cv.cvtColor(image, cv.COLOR_BGR2RGB) # 转Tensor、调整维度顺序、转float、归一化 input_tensor = torch.from_numpy(image_rgb).permute(2, 0, 1).to(torch.float32) / 255.0 color_jitter = torchvision.transforms.ColorJitter(brightness=0.5, hue=0.3) jitted_img_list = [color_jitter(input_tensor) for _ in range(4)] # 逐张保存增强后的图像 for img_idx, aug_tensor in enumerate(jitted_img_list): # 维度转回HWC、转numpy数组、还原像素值域、转uint8、通道顺序转回BGR aug_img_np = aug_tensor.permute(1, 2, 0).numpy() aug_img_np = (aug_img_np * 255).astype(np.uint8) aug_img_bgr = cv.cvtColor(aug_img_np, cv.COLOR_RGB2BGR) cv.imwrite(f"jitted_result_{img_idx}.png", aug_img_bgr)
如果你需要把4张增强后的图像拼成一张大图再保存,可以在保存前用
np.concatenate或np.hstack/np.vstack把4张numpy数组拼合后,再调用cv2.imwrite写入即可。
内容的提问来源于stack exchange,提问作者Interpreter67
相关产品推荐
相关产品推荐

