如何为分割掩码第二维度添加虚拟维度适配torch.utils.make_grid?
解决分割掩码维度匹配问题,适配TensorBoardX的
make_grid 嘿,这个问题我之前也碰到过!咱们来一步步解决它——你需要把形状为[5,1,100,100]的分割掩码扩展为[5,3,100,100],和RGB图像维度对齐,关键是选对维度扩展的方法,我帮你拆解一下:
为什么你之前的方法没成功?
view失败:view要求张量的总元素数不变,原掩码总元素是5*1*100*100=50000,而目标形状是5*3*100*100=150000,元素数不匹配,所以必然报错。unsqueeze没用:你的掩码第二维度已经是1了,再unsqueeze会增加新维度(比如变成[5,1,1,100,100]),反而离目标更远。expand用错参数:如果没指定对维度的扩展次数,或者没保持其他维度不变,也达不到效果。
两种可行的解决方案
方案1:用repeat复制通道维度(会占用额外内存)
repeat会直接复制数据,把第二维度的单个通道复制3次,完美匹配RGB图像的通道数:
import torch # 假设你的分割掩码是mask,形状[5,1,100,100] mask = torch.randint(0, 2, (5, 1, 100, 100)) # 示例掩码,0/1分割值 mask_3ch = mask.repeat(1, 3, 1, 1) # 参数对应:[batch维度重复次数, 通道维度重复次数, 高度维度重复次数, 宽度维度重复次数] print(mask_3ch.shape) # 输出: torch.Size([5, 3, 100, 100])
方案2:用expand创建视图(不占用额外内存)
expand不会复制数据,只是创建一个新的视图,适合不想增加内存开销的场景(前提是原通道维度为1):
mask_3ch = mask.expand(-1, 3, -1, -1) # -1表示保持对应维度的原大小 print(mask_3ch.shape) # 输出: torch.Size([5, 3, 100, 100])
和RGB图像一起在TensorBoardX展示的示例
搞定维度后,就可以和RGB图像一起用make_grid生成网格图,然后传入TensorBoard了:
from torch.utils.tensorboard import SummaryWriter from torchvision.utils import make_grid # 假设你的RGB图像批次是images,形状[5,3,100,100] images = torch.randn(5, 3, 100, 100) # 拼接图像和掩码(这里选择在batch维度拼接,方便并排显示) combined = torch.cat([images, mask_3ch], dim=0) # 生成网格,nrow设为5,这样一行显示一个batch的图像和对应的掩码 grid = make_grid(combined, nrow=5, padding=2) # 写入TensorBoard writer = SummaryWriter(log_dir='./logs') writer.add_image('Original Images vs Segmentation Masks', grid, global_step=0) writer.close()
内容的提问来源于stack exchange,提问作者rmm
相关产品推荐
相关产品推荐

