PyTorch GPU并行计算优化:torchdiff求解ODE的torch.cat性能瓶颈
问题描述
我正在使用torchdiff求解常微分方程组的初值问题,forward函数的输入y0是尺寸为120000的一维张量。在函数中,我先将其切分为12个长度为10000的一维张量A、B、C……L。此外还有若干系数k1、k2、k3……,并行计算中被迫大量使用torch.cat操作,导致GPU利用率较低。请问是否有算法优化方式提升计算速度?
当前实现代码:
N_SP = 10000 zeros_mat = torch.zeros(N_SP) def forward(self, t, y0, rates, ba_input): k1, \ k2, k3, k4\ k5, k6... = rates A = y0[0:N_sp] B = y0[N_sp:2 * N_sp] C = y0[2 * N_sp:3 * N_sp] D = y0[3 * N_sp:4 * N_sp] E = y0[4 * N_sp:5 * N_sp] F = y0[5 * N_sp:6 * N_sp] G = y0[6 * N_sp:7 * N_sp] H = y0[7 * N_sp:8 * N_sp] I = y0[8 * N_sp:9 * N_sp] J = y0[9 * N_sp:10 * N_sp] K = y0[10 * N_sp:11 * N_sp] L = y0[11 * N_sp:12 * N_sp] y1 = torch.cat((k1*A,k5*B,k5*B,k2*D,k7*G,k4*G,zeros_mat,zeros_mat,zeros_mat,k2*H,k3*I,k4*J),0) *torch.cat((L,L,L,L,L,L,zeros_mat,zeros_mat,zeros_mat,K,K,L)) + \ torch.cat((k7*D,k4*A,K4*J,K6*B,K3*B,K2*E,K3*A,K6*H,K9*I,K3*H,torch.ones(N_sp),ba_input),0)+ torch.cat((zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,L**4,L**4),0)* torch.cat((zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,zeros_mat,H,H),0) return y1
补充说明:这是使用Torchdiff库的odeint函数求解的高耦合常微分方程组,若无需张量拼接,GPU利用率可接近100%。此外,曾尝试在外部初始化全局张量数组,通过切片操作填充所需“块”而非拼接,但效果不佳。
优化方案
重塑张量维度,简化切片操作
直接将一维的y0重塑为(12, N_SP)的二维张量,通过unbind一次性拆分出A-L,比逐个切片更高效:y_reshaped = y0.view(12, N_SP) A, B, C, D, E, F, G, H, I, J, K, L = y_reshaped.unbind(dim=0)预分配结果张量,按块赋值替代torch.cat
这是提升GPU利用率的核心操作——预先创建和y0同形状的空张量,直接给对应位置赋值,完全避免频繁拼接带来的内存拷贝和同步开销:def forward(self, t, y0, rates, ba_input): k1, k2, k3, k4, k5, k6, k7, k9 = rates N_SP = 10000 # 重塑拆分y0 y_reshaped = y0.view(12, N_SP) A, B, C, D, E, F, G, H, I, J, K, L = y_reshaped.unbind(dim=0) # 预分配各部分计算张量 part1 = torch.zeros(12, N_SP, device=y0.device) part2 = torch.zeros(12, N_SP, device=y0.device) part3 = torch.zeros(12, N_SP, device=y0.device) # 填充第一部分乘积项 part1[0] = k1 * A * L part1[1] = k5 * B * L part1[2] = k5 * B * L part1[3] = k2 * D * L part1[4] = k7 * G * L part1[5] = k4 * G * L part1[9] = k2 * H * K part1[10] = k3 * I * K part1[11] = k4 * J * L # 填充第二部分加项 part2[0] = k7 * D part2[1] = k4 * A part2[2] = k4 * J # 修正原代码中K4的笔误 part2[3] = k6 * B part2[4] = k3 * B part2[5] = k2 * E part2[6] = k3 * A part2[7] = k6 * H part2[8] = k9 * I part2[9] = k3 * H part2[10] = self.ones_mat # 预创建的常量张量 part2[11] = ba_input # 填充第三部分乘积项 L_pow4 = L ** 4 # 复用计算结果 part3[10] = L_pow4 * H part3[11] = L_pow4 * H # 合并并展平为一维张量返回 y1 = (part1 + part2 + part3).flatten() return y1预创建GPU常量张量
在类初始化阶段提前创建并移至GPU,避免每次forward重复创建:def __init__(self): N_SP = 10000 self.zeros_mat = torch.zeros(N_SP, device='cuda') self.ones_mat = torch.ones(N_SP, device='cuda')复用重复计算结果
对原代码中重复出现的L**4这类计算,提前计算一次后复用,减少冗余运算。
内容的提问来源于stack exchange,提问作者Edler
相关产品推荐
相关产品推荐

