PyTorch自动编码器模型摘要MPS报错及维度转换疑问
问题解决与疑问解答
一、MPS设备上torchsummary报错的解决办法
你遇到的RuntimeError: Placeholder storage has not been allocated on MPS device!,是因为torchsummary默认会在CPU上生成输入占位张量,推送到MPS设备时出现兼容性问题。解决方法如下:
- 手动创建符合输入维度的张量,明确指定放在MPS设备上,传给summary函数的
input_data参数,替代input_size。 - 示例代码:
import torch from torchsummary import summary # 假设模型已部署到MPS设备 model = model.to('mps') # 创建匹配MNIST输入维度的张量(单通道28x28,batch_size设为1即可) input_tensor = torch.randn(1, 1, 28, 28).to('mps') # 传入input_data获取模型摘要 summary(model, input_data=input_tensor)
二、28x28转784维度的实现说明
这个维度转换不是框架自动完成的,需要你在模型的forward方法里手动实现,常见两种写法:
- 用
view方法:x = x.view(x.size(0), -1),其中-1会自动计算出784(28*28),x.size(0)保留batch维度。 - 用
flatten方法:x = torch.flatten(x, start_dim=1),start_dim=1表示从通道维度后开始展平,跳过batch维度。 - 该操作一般放在编码器的最开始,把二维图像张量转成一维向量,才能喂给后续全连接层做编码。
内容的提问来源于stack exchange,提问作者user1357015
相关产品推荐
相关产品推荐

