PyTorch vmap批量转PIL图像遇RuntimeError及维度错误问题
批量转换分割模型预测结果为Pillow图像的问题解决
问题回顾
在CPU环境下将分割模型的批量预测结果(形状(batch_size, 1, H, W))转换为Pillow图像时:
- 直接使用
ToPILImage转换4维张量,触发维度错误:ValueError: pic should be 2/3 dimensional. Got 4 dimensions. - 尝试用
torch.func.vmap包装ToPILImage实现批量转换,触发存储访问错误:RuntimeError: Cannot access data pointer of Tensor that doesn't have storage
错误原因
- 维度错误原因:
ToPILImage仅接受2维(H, W)或3维(C, H, W)张量,带batch维度的4维张量不符合要求。 - vmap报错原因:
vmap是为纯张量运算设计的自动向量化工具,而ToPILImage内部需要直接访问张量的底层数据指针(data_ptr())。vmap处理时会生成无实际存储的虚拟张量,导致无法获取数据指针,因此不能用vmap包装ToPILImage。
正确实现方式
无需使用vmap,直接拆分batch后逐个转换即可,以下是两种简洁实现:
方式1:遍历batch样本转换
import torch from torchvision.transforms import ToPILImage transform = ToPILImage() model.eval() with torch.no_grad(): for i, (x, y) in enumerate(dataloader): y_hat = torch.sigmoid(model(x)) y_hat = (y_hat > 0.5).float() img_batch = [] # 遍历每个样本,去除单通道维度后转换 for pred in y_hat: img = transform(pred.squeeze(0)) # 将(1, H, W)转为(H, W) img_batch.append(img)
方式2:列表推导式+unstack简化
import torch from torchvision.transforms import ToPILImage transform = ToPILImage() model.eval() with torch.no_grad(): for i, (x, y) in enumerate(dataloader): y_hat = torch.sigmoid(model(x)) y_hat = (y_hat > 0.5).float() # unstack拆分batch维度,逐个处理 img_batch = [transform(pred.squeeze(0)) for pred in torch.unbind(y_hat)]
补充说明
- 若分割结果是多通道(如
(batch_size, 3, H, W)),则无需squeeze(0),直接传入(C, H, W)形状的张量给ToPILImage即可。 - 推理阶段添加
torch.no_grad()可禁用梯度计算,降低CPU内存消耗。
内容的提问来源于stack exchange,提问作者codersblock
相关产品推荐
相关产品推荐

