PyTorch 0.17版本中如何从张量保存.exr文件?
解决torchvision.save_image保存EXR文件报错的问题
问题原因
torchvision 0.17及后续版本移除了save_float_image函数,剩余的save_image仅支持PNG、JPG等常规图像格式,原生不支持EXR格式,因此直接写入EXR文件会触发“unknown file extension *.exr”错误。
可行解决方案
方案1:通过OpenCV保存EXR
OpenCV支持读写EXR格式(需确保安装的OpenCV版本包含OpenEXR支持),步骤如下:
- 将PyTorch张量从CHW维度转为HWC,再转换为NumPy数组
- 直接用OpenCV的
imwrite写入EXR文件
示例代码:
import cv2 import numpy as np import torch import os # 示例输入张量:3通道、256x256的浮点张量 input_tensor = torch.randn(3, 256, 256, dtype=torch.float32) # 张量转NumPy并调整维度顺序 img_array = input_tensor.permute(1, 2, 0).cpu().numpy() # 保存EXR文件 save_path = os.path.join('folder/', 'foo.exr') cv2.imwrite(save_path, img_array)
方案2:使用pyexr库专门处理
pyexr是针对EXR格式的专用库,支持直接写入浮点图像数据:
- 先安装依赖:
pip install pyexr - 转换张量为NumPy数组后直接写入
示例代码:
import pyexr import torch import os input_tensor = torch.randn(3, 256, 256, dtype=torch.float32) img_array = input_tensor.permute(1, 2, 0).cpu().numpy() save_path = os.path.join('folder/', 'foo.exr') pyexr.write(save_path, img_array)
注意事项
- EXR用于存储高精度浮点图像,确保输入张量为
float32或float64类型 - 若张量在GPU上,需先调用
.cpu()转移到CPU再转换为NumPy数组
内容的提问来源于stack exchange,提问作者0xbadf00d
相关产品推荐
相关产品推荐

