PyTorch中图像变换与数据增强代码运行异常问题求助
PyTorch图像增强代码运行故障排查与修正
现有代码的核心问题
- 未传入待处理图像:定义完
augmentation增强序列后,没有将读取到的image变量传入执行计算,直接把Sequential模块对象传入cv2.imwrite,参数类型完全不匹配,无法执行写入。 - 数据格式不兼容:OpenCV通过
cv2.imread读取的图像是BGR通道顺序、HWC维度排列的numpy数组,而torchvision transforms的nn.Sequential封装形式默认接收形状为[C, H, W]、数值范围在[0,1]的浮点Tensor,直接传入numpy数组会触发类型/维度校验报错。 - 命名空间拼写错误:代码中读取图像用的是
cv2.imread,写入时写的是cv.imwrite,导入别名前后不一致会触发命名不存在的报错。 - 输出格式不匹配:OpenCV的
imwrite方法仅支持接收HWC维度排列、BGR通道顺序的numpy数组格式图像,无法直接写入PyTorch Tensor类型数据。
修正后可运行代码
import cv2 import torch import torchvision from torchvision import transforms from PIL import Image import numpy as np # 读取原始图像 image = cv2.imread('image.png') # 转换通道顺序为RGB,再转为PIL Image适配transforms输入要求 image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) image_pil = Image.fromarray(image_rgb) # 组装完整处理流程:先转Tensor,再执行自定义增强 aug_pipeline = transforms.Compose([ transforms.ToTensor(), torch.nn.Sequential( transforms.RandomGrayscale(p=0.5), transforms.RandomVerticalFlip(p=0.5) ) ]) # 传入图像执行增强 augmented_tensor = aug_pipeline(image_pil) # 将增强后的Tensor转换为OpenCV兼容的格式 # 维度从CHW转为HWC,数值范围从[0,1]转回[0,255]的uint8类型,通道从RGB转回BGR augmented_np = augmented_tensor.permute(1, 2, 0).numpy() * 255 augmented_np = augmented_np.astype(np.uint8) augmented_bgr = cv2.cvtColor(augmented_np, cv2.COLOR_RGB2BGR) # 保存增强后的图像 cv2.imwrite('test.png', augmented_bgr)
补充说明
如果不想引入PIL做格式中转,可以直接手动完成numpy数组到Tensor的维度、通道、数值范围转换,但必须保证输入transforms的数据格式符合要求;增强完成后必须转回OpenCV支持的numpy数组格式,才能正常调用cv2.imwrite保存。
内容的提问来源于stack exchange,提问作者Interpreter67
相关产品推荐
相关产品推荐

