PyTorch自定义DCGAN中选择性高斯模糊致反向传播过慢求助
解决DCGAN中自定义高斯模糊导致
backward()耗时激增的问题 问题背景
用PyTorch构建自定义DCGAN时,在生成器末尾添加了一个仅对超过特定阈值的像素执行高斯模糊的滤波器。添加后,backward()调用耗时从几乎瞬间飙升至超过一分钟,推测是遍历图像的嵌套循环导致了性能瓶颈。
原生成器前向传播代码:
def forward(self, x): x = self.gen(x) x = convolve2D(x) return x
原自定义卷积函数(性能瓶颈根源):
def convolve2D(batch, padding=0, strides=1): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") kernel = torch.tensor(([1, 2, 1], [2, 4, 2], [1, 2, 1])).to(device) kernel_sum = kernel.sum() # Gather Shapes of Kernel + Image + Padding xKernShape = kernel.shape[0] yKernShape = kernel.shape[1] xImgShape = batch.shape[2] yImgShape = batch.shape[3] copy = batch.clone() for i, image in tqdm(enumerate(copy)): for j, channel in enumerate(image): # Apply Equal Padding to All Sides if padding != 0: channelPadded = torch.zeros((channel.shape[0] + padding*2, channel.shape[1] + padding*2)) channelPadded[int(padding):int(-1 * padding), int(padding):int(-1 * padding)] = channel print(channelPadded) else: channelPadded = channel # Iterate through image for y in range(yImgShape): # Exit Convolution if y > yImgShape - yKernShape: break # Only Convolve if y has gone down by the specified Strides if y % strides == 0: for x in range(xImgShape): # Go to next row once kernel is out of bounds if x > xImgShape - xKernShape: break # Ignore if pixel is an edge if channel[x + 1, y + 1] < 0: continue else: # Only Convolve if x has moved by the specified Strides if x % strides == 0: batch[i][j][x + 1, y + 1] = torch.mul(kernel, channelPadded[x: x + xKernShape, y: y + yKernShape]).sum() / kernel_sum return batch
问题分析
这段代码的核心问题是四层Python嵌套循环:遍历batch、通道、图像的x/y轴。PyTorch的自动求导机制会追踪每一次循环里的张量操作,生成大量细碎的计算图节点,反向传播时需要逐个处理这些节点,直接导致耗时爆炸。另外手动处理padding、逐像素赋值的操作完全放弃了PyTorch的GPU向量化计算优势,进一步拖慢了速度。
优化方案
用PyTorch内置的向量化卷积函数替代Python循环,步骤如下:
- 对整个batch做全图高斯模糊,用
torch.nn.functional.conv2d实现,利用GPU并行计算。 - 生成掩码:标记出需要替换成模糊结果的像素(即原图像中像素值≥0的位置,对应原代码里不跳过的像素)。
- 用掩码合并原图像和模糊后的图像:仅在掩码指定的位置替换为模糊值,其余位置保留原像素。
优化后的代码
import torch.nn.functional as F def selective_gaussian_blur(batch, threshold=0.0, padding=1, stride=1): device = batch.device # 直接复用输入张量的设备,避免重复判断 # 定义高斯核,转换成conv2d要求的格式: (out_channels, in_channels//groups, kernel_h, kernel_w) kernel = torch.tensor([[1, 2, 1], [2, 4, 2], [1, 2, 1]], dtype=batch.dtype, device=device) kernel = kernel / kernel.sum() # 提前完成归一化 # 适配多通道场景:用分组卷积保证每个通道单独模糊 in_channels = batch.shape[1] kernel = kernel.view(1, 1, 3, 3).repeat(in_channels, 1, 1, 1) # 对整个batch执行高斯模糊 blurred_batch = F.conv2d(batch, kernel, padding=padding, stride=stride, groups=in_channels) # 生成掩码:仅保留原图像中像素值≥阈值的位置(对应原代码不跳过的逻辑) mask = (batch >= threshold).float() # 合并原图像与模糊图像:掩码区域用模糊值,其余保留原值 result = mask * blurred_batch + (1 - mask) * batch return result # 生成器前向传播修改为 def forward(self, x): x = self.gen(x) x = selective_gaussian_blur(x) return x
优化说明
- 完全移除Python循环,所有操作都是PyTorch的向量化GPU运算,前向传播速度大幅提升,同时自动求导的计算图会非常简洁,反向传播耗时回归正常水平。
- 用
groups=in_channels实现分组卷积,保证每个通道单独处理,和原代码逻辑完全一致。 - 直接复用输入张量的设备,避免不必要的张量跨设备移动。
- 掩码合并操作精准保留了“仅对超过阈值的像素模糊”的核心需求。
内容的提问来源于stack exchange,提问作者mattwaz
相关产品推荐
相关产品推荐

