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

如何加速稀疏张量下的torch.einsum批量矩阵乘法运算?

稀疏结构下的批量矩阵乘法加速方案

针对你提到的proj = torch.einsum('abi,aic->abc', A, B)运算在n=50k时速度极慢的问题,由于A的前两维是稀疏结构,我们可以通过只计算非零元素对应的运算来大幅降低计算量,具体实现步骤如下:

核心思路拆解

原运算的本质是:对每个索引对(i,j),proj[i,j,:] = A[i,j,:] @ B[i,:,:](即1xd向量与dxd矩阵的乘法)。全量计算会对大量零向量做无用运算,我们只需聚焦A的非零元素对应的(i,j)对,跳过零值计算。

具体实现步骤

假设你已经将稀疏的A表示为以下两个张量:

  • indices: 形状为[K, 2]的长整型张量,存储所有非零元素的索引对(i,j),K是非零元素的总数量
  • values_A: 形状为[K, d]的张量,存储每个(i,j)对应的1xd非零向量
  1. 提取对应B的切片
    根据indices中的i索引,从B中取出对应的dxd矩阵切片:

    B_slices = B[indices[:, 0]]  # 形状:[K, d, d]
    
  2. 计算非零元素对应的proj值
    将values_A扩展维度后与B_slices做批量矩阵乘法,得到每个(i,j)对应的proj结果:

    # 将values_A从[K,d]转为[K,1,d],与B_slices做批量矩阵乘法后压缩维度
    proj_values = torch.bmm(values_A.unsqueeze(1), B_slices).squeeze(1)  # 形状:[K, d]
    
  3. 组装全量结果(可选)
    如果需要得到与原运算一致的全量nxnxd张量,初始化零张量后将计算结果填入对应位置:

    proj = torch.zeros(n, n, d, device=values_A.device, dtype=values_A.dtype)
    proj[indices[:, 0], indices[:, 1]] = proj_values
    

效率提升说明

原全量计算的时间复杂度为O(n²*d²),当n=50k时n²=2.5e9,计算量极其庞大。优化后的复杂度为O(K*d²),只要K远小于n²(稀疏场景下通常如此),计算量会呈数量级下降,速度提升非常明显。

额外优化建议

  • 确保所有张量在同一设备(CPU/GPU)上运行,避免跨设备数据传输的开销
  • 如果不需要全量的proj张量,可以直接保留稀疏格式的indices和proj_values,进一步节省内存和计算资源

内容的提问来源于stack exchange,提问作者Neurobro

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 04:30:59