如何在PyTorch中高效为PIL图像添加边框增强?
嘿,我完全懂你的困扰——频繁在PIL和Tensor之间来回转换确实是性能杀手,而且你原来的黑色圆形边框用双重循环计算像素的方式也太拖速度了。我来给你重构一下代码,把这两个核心瓶颈解决掉,让你的数据增强速度快起来:
优化后的边框增强Transformer
下面的代码从减少格式转换、用向量化操作替代循环、缓存计算结果三个方面做了优化,同时还修正了原代码里的概率逻辑偏差:
import torch from torchvision import transforms from PIL import Image class FastBorderTransformer(torch.nn.Module): def __init__(self, black_p=0.05, white_p=0.01, white_border_prop=0.05): super().__init__() self.white_p = white_p self.black_p = black_p assert white_border_prop < 0.5, "Can't add white border bigger than full image" self.white_border_prop = white_border_prop # 缓存不同尺寸的黑色圆形mask,避免重复计算 self.black_mask_cache = {} self.to_tensor = transforms.ToTensor() self.to_pil = transforms.ToPILImage() def white_square_border(self, img_tensor): """直接在Tensor上完成边框绘制,省去来回转换开销""" _, width, height = img_tensor.size() h_border = int(self.white_border_prop * width) v_border = int(self.white_border_prop * height) # 用Tensor切片一次性设置所有边框区域,比多次赋值更高效 img_tensor[:, :h_border, :] = 1.0 img_tensor[:, width - h_border:, :] = 1.0 img_tensor[:, :, :v_border] = 1.0 img_tensor[:, :, height - v_border:] = 1.0 return img_tensor def _get_black_circle_mask(self, dim): """用向量化操作生成圆形mask,替代原有的双重循环""" if dim not in self.black_mask_cache: # 生成坐标网格,一次性计算所有像素的位置 y, x = torch.meshgrid(torch.arange(dim), torch.arange(dim), indexing='ij') center = dim // 2 # 计算每个像素到中心的距离平方 dist_sq = (x - center)**2 + (y - center)**2 # 标记圆外需要涂黑的区域 mask = dist_sq > (dim**2) # 扩展为3通道,方便直接和图像Tensor匹配 mask = mask.unsqueeze(0).repeat(3, 1, 1) self.black_mask_cache[dim] = mask return self.black_mask_cache[dim] def black_circle_border(self, img_tensor): """直接在Tensor上应用预生成的mask""" _, width, height = img_tensor.size() # 如果你的图像不是正方形,可以改成dim = min(width, height),再调整mask的位置 assert width == height, "Black circle border requires square image, adjust logic if needed" mask = self._get_black_circle_mask(width).to(img_tensor.device) img_tensor[mask] = 0.0 return img_tensor def forward(self, img): """只做一次PIL→Tensor转换,所有操作完成后再转回PIL""" # 仅转换一次Tensor,避免多次转换的开销 img_tensor = self.to_tensor(img) # 一次随机采样判断操作类型,修正原代码的概率偏差 rand_val = torch.rand(1) if rand_val < self.black_p: img_tensor = self.black_circle_border(img_tensor) elif rand_val < self.black_p + self.white_p: img_tensor = self.white_square_border(img_tensor) # 最后仅转回一次PIL return self.to_pil(img_tensor)
核心优化点说明
减少格式转换次数
原代码在每个边框操作里都做PIL→Tensor→PIL的来回转换,现在只在forward开头转一次Tensor,所有增强操作完成后再转一次PIL,直接砍掉了大部分转换开销。向量化替代双重循环
黑色圆形边框的计算原来用嵌套循环遍历每个像素,现在用torch.meshgrid生成坐标网格,一次性完成所有像素的距离计算,速度能提升几个数量级。缓存mask避免重复计算
针对不同尺寸的图像缓存对应的圆形mask,不用每次处理相同尺寸的图像都重新计算,进一步节省计算时间。修正概率逻辑偏差
原代码先判断黑色边框概率,不满足时再重新采样判断白色边框,这会导致白色边框的实际概率是(1-black_p)*white_p,和设定值不符。现在用一次随机采样判断区间,保证概率完全符合你的配置。
如果你的数据加载流程里已经把图像转成了Tensor,还可以直接去掉to_tensor和to_pil的转换步骤,速度还能再上一个台阶。
内容的提问来源于stack exchange,提问作者shmulvad
相关产品推荐
相关产品推荐

