如何加速稀疏张量下的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非零向量
提取对应B的切片
根据indices中的i索引,从B中取出对应的dxd矩阵切片:B_slices = B[indices[:, 0]] # 形状:[K, d, d]计算非零元素对应的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]组装全量结果(可选)
如果需要得到与原运算一致的全量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
相关产品推荐
相关产品推荐

