Torchvision Resize无法识别PyTorch二维张量维度问题求助
问题详情
- 原始张量形状为
[512, 512],尝试用torchvision.transforms.Resize调整为[256, 256],代码如下:resized = T.Resize(size=(256,256))(img) - 首次报错:
Input and output must have the same number of spatial dimensions, but got input with spatial dimensions of [512] and output size of [256, 256]. Please provide input tensor in (N, C, d1, d2, ...,dK) format and output size in (o1, o2, ...,oK) format.
- 截取张量
img = img[:100,:200]后,报错变为:Input and output must have the same number of spatial dimensions, but got input with spatial dimensions of [200] and output size of [256, 256]. Please provide input tensor in (N, C, d1, d2, ...,dK) format and output size in (o1, o2, ...,oK) format.
原因分析
torchvision.transforms.Resize要求输入张量必须包含通道维度,格式为(C, H, W)(单张图片)或(N, C, H, W)(批量图片)。而原始张量是纯二维的(H, W),Resize会误将其识别为1维空间输入(仅读取最后一个维度作为空间维度),导致维度不匹配报错。
解决方法
方法1:添加通道维度后使用Resize
先给二维张量插入通道维度,处理完成后再移除多余维度:
import torchvision.transforms as T # 为二维张量添加通道维度,变为(1, 512, 512) img_with_channel = img.unsqueeze(0) # 执行Resize resized_tensor = T.Resize(size=(256, 256))(img_with_channel) # 移除通道维度,变回(256, 256) resized_tensor = resized_tensor.squeeze(0)
方法2:使用torch.nn.functional.interpolate
这个方法对输入维度兼容性更强,无需额外调整也能实现缩放(需要临时添加批量和通道维度):
import torch.nn.functional as F # 添加批量和通道维度,变为(1, 1, 512, 512) resized_tensor = F.interpolate(img[None, None, ...], size=(256, 256)) # 移除多余维度,得到(256, 256) resized_tensor = resized_tensor.squeeze()
内容的提问来源于stack exchange,提问作者Iliasp

