PyTorch自定义函数反向传播耗时增至4倍:原因与优化方案咨询
反向传播耗时激增的原因分析与代码优化方案
耗时激增的核心原因
- 高维度批量运算的计算量过载:你的输入M维度是
16×10000×2×2,意味着要对10000个空间点分别执行2×2矩阵乘法,再和21×21的Z向量做三次矩阵乘法运算。单批次浮点运算量(FLOPs)极高,反向传播时还要为这些矩阵运算、exp函数、除法操作计算梯度,进一步放大了耗时。 - 冗余的维度变换与内存访问低效:多次
unsqueeze、permute操作会导致张量内存布局不连续,GPU无法高效利用内存带宽,反向传播时还要处理这些维度变换的梯度,额外增加开销。 - 归一化步骤的额外计算:手动执行
exp → sum → 除法的归一化流程,反向传播时需要分别计算sum和除法的梯度,相比PyTorch优化后的内置函数,这部分的计算效率更低。
针对性优化方案
1. 预计算固定张量,避免重复创建
Z张量仅和kernel_size相关,训练阶段无需每次调用函数都重新生成,提前预计算并放到GPU上,可节省大量重复创建张量的时间。
2. 用爱因斯坦求和(einsum)简化矩阵运算
Z^T @ INV_SIGMA @ Z是典型的二次型运算,用torch.einsum可以直接表达运算逻辑,替代多次matmul和permute,减少维度变换的同时,让PyTorch自动优化运算流程,提升GPU利用率。
3. 用内置softmax替代手动归一化
你当前的归一化逻辑raw_kernel / sum(raw_kernel)等价于softmax(-0.5 * quadratic_form, dim=(-2,-1))。PyTorch的nn.functional.softmax经过深度优化,无论是前向还是反向传播,效率都远高于手动实现的除法归一化。
4. 简化维度变换操作
对最后两个维度的转置,用transpose(-1, -2)替代permute(0,1,3,2),语义更清晰且运算更高效。
5. 可选:启用混合精度训练
V100支持FP16混合精度,通过torch.cuda.amp可以在不损失精度的前提下,大幅降低内存占用并提升计算速度,尤其适合高维度张量运算场景。
优化后的代码示例
# 预计算Z(训练前执行一次即可,kernel_size固定时无需重复生成) kernel_size = 21 # 使用indexing='ij'确保meshgrid的维度符合预期(PyTorch 1.10+支持) X, Y = torch.meshgrid(torch.arange(kernel_size), torch.arange(kernel_size), indexing='ij') Z = torch.stack((Y, X), dim=-1).cuda() # 直接放到GPU上 def optimized_batch_function(M, Z, sf=2): ''' Input: M: b x (h*w) x 2 x 2 torch tensor Z: k x k x 2 precomputed tensor (same device as M) sf: scale factor Output: kernel: b x (h*w) x k x k torch tensor ''' # 计算转置并得到INV_SIGMA,替代permute更高效 M_t = M.transpose(-1, -2) INV_SIGMA = torch.matmul(M_t, M) # b x (h*w) x 2 x 2 # 用einsum直接计算二次型,省去多次维度变换和matmul quadratic_form = torch.einsum('bnij, kli, klj -> bnkl', INV_SIGMA, Z, Z) # 用softmax完成归一化,前向+反向效率更高 x = -0.5 * quadratic_form kernel = torch.nn.functional.softmax(x, dim=(-2, -1)) return kernel
额外优化建议
- 可以用
torch.jit.script编译优化后的函数,进一步消除Python overhead,提升运行速度:optimized_batch_function = torch.jit.script(optimized_batch_function) - 若
kernel_size需要动态调整,可以将Z的预计算逻辑封装成一个辅助函数,在每次kernel_size变化时重新生成即可。
内容的提问来源于stack exchange,提问作者Yuji Lin
相关产品推荐
相关产品推荐

