PyTorch1.7.1与1.9.1版本中torch.utils.save_image的差异问题咨询
PyTorch(torchvision)不同版本
save_image函数的功能差异说明 首先明确:save_image是torchvision.utils下的工具函数,功能逻辑会跟随torchvision版本迭代变动,和你使用的PyTorch版本是对应绑定的,你遇到的报错确实是版本差异导致的。
两个版本的具体差异
- PyTorch 1.7.1 对应绑定的 torchvision 版本为 0.8.x:该版本的
save_image没有对输入张量的通道数做强校验,只要输入张量的维度格式符合[C, H, W]或者[N, C, H, W],就会直接传递给底层PIL接口处理,PIL自动识别4通道张量保存为RGBA格式的图片,所以你的代码可以正常运行。 - PyTorch 1.9.1 对应绑定的 torchvision 版本为 0.10.x:该版本在
save_image的预处理逻辑中新增了默认通道匹配规则,默认仅支持1通道(灰度)、3通道(RGB)输入,在处理padding、范围映射等逻辑时会硬编码匹配3通道参数,输入4通道张量时就会触发你看到的张量a(3)和张量b(4)维度不匹配的报错。
可替代解决方案
除了你目前使用的先转RGB再保存的方案外,还有两种可选方案:
- 保留RGBA通道的前提下,直接走PIL保存逻辑:
from PIL import Image import numpy as np # 假设final_images是[4,256,256]的cuda张量,数值范围0-1 img_np = (final_images.cpu().permute(1,2,0).numpy() * 255).astype(np.uint8) Image.fromarray(img_np, mode='RGBA').save("./output/results_page_{}.png".format(count))
- 升级torchvision到0.11.0及以上版本:该版本官方修复了4通道支持的问题,重新兼容RGBA张量的直接保存,你的原有代码可以正常运行。
内容的提问来源于stack exchange,提问作者Bluce.Q
相关产品推荐
相关产品推荐

