如何使用PyTorch的F.interpolate将MNIST图像张量从[1,28,28]缩放到[1,14,14]?
解决PyTorch中MNIST图像张量缩放的问题
嘿,我来帮你搞定这个图像缩放的问题!你遇到的问题其实是因为F.interpolate对输入张量的维度格式有要求,咱们一步步来拆解:
问题根源
你用的images[0]形状是torch.Size([1,28,28]),也就是**(通道数C, 高度H, 宽度W),但F.interpolate默认期望的输入格式是(批量数B, 通道数C, 高度H, 宽度W)**——简单说就是少了一个「批量维度」。这就导致:
- 当你传
F.interpolate(images[0], 14)时,函数把输入当成了(C, H)格式,只缩放了最后一个维度,所以得到[1,28,14]; - 当你传
(14,14)时,函数误以为输入是1D序列,而你给了2个尺寸参数,自然触发维度不匹配的错误。
正确实现方法
方法1:给单个图像添加批量维度
只需要用unsqueeze(0)给张量加一个批量维度,处理完后再按需去掉即可:
# 给单个图像添加批量维度,形状变为 [1, 1, 28, 28] img_with_batch = images[0].unsqueeze(0) # 执行插值缩放,指定目标尺寸(14,14) resized_img = F.interpolate(img_with_batch, size=(14, 14)) # 去掉批量维度,得到最终目标形状 [1,14,14] resized_img = resized_img.squeeze(0) # 验证形状 print(resized_img.shape) # 输出: torch.Size([1, 14, 14])
方法2:直接对批量张量处理(更高效)
如果你是处理整个批次的图像,不需要单独取images[0],直接对images(形状[B,1,28,28])做插值即可:
# 对整个批次缩放,得到形状 [B,1,14,14] resized_images = F.interpolate(images, size=(14,14)) # 取第一个图像,形状就是 [1,14,14] print(resized_images[0].shape)
这样就能完美得到你想要的[1,14,14]形状的张量啦!
内容的提问来源于stack exchange,提问作者Data Mastery
相关产品推荐
相关产品推荐

