You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 22:41:00