You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch vmap批量转PIL图像遇RuntimeError及维度错误问题

批量转换分割模型预测结果为Pillow图像的问题解决

问题回顾

在CPU环境下将分割模型的批量预测结果(形状(batch_size, 1, H, W))转换为Pillow图像时:

  1. 直接使用ToPILImage转换4维张量,触发维度错误:
    ValueError: pic should be 2/3 dimensional. Got 4 dimensions.
    
  2. 尝试用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.26 12:15:15