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

PyTorch两种向量化距离计算代码的性能差异原因咨询

PyTorch 成对平方欧氏距离实现的性能问题

问题背景

给定形状为(M, d)的矩阵A和形状为(N, d)的矩阵B,需要计算所有行对的欧氏距离平方矩阵D,满足D[i,j] = torch.sum((A[i] - B[j])**2)。
目前有两种无显式循环的向量化实现,计算结果完全一致,但性能差距可达100倍:

高性能矩阵乘法实现(基于恒等式 $||v-w||^2 = ||v||^2 + ||w||^2 - 2v·w$)

dists = (torch.sum(torch.square(A),dim=1).view((-1,1)) 
          + torch.sum(torch.square(B),dim=1).view((1,-1))
          - 2*A @ B.t())

低性能广播实现

A_v = A.view((A.shape[0],-1,1))
B_v = B.view((B.shape[0],-1)).permute((1,0)) 
dists=torch.sum(torch.square(A_v-B_v),dim=1)

为什么广播实现在GPU上效率远低于矩阵乘法实现

  • 中间张量显存开销天差地别:广播实现执行A_v - B_v时,会通过广播生成形状为(M, d, N)的三维中间张量。举个常见场景:当M=N=1024、特征维度d=512时,这个三维张量的元素总数是5.36亿,单精度浮点下就占2GB显存;而矩阵乘法实现全程最大的张量就是最终输出的(M,N)距离矩阵(仅104万元素,占4MB显存),加上两个长度分别为M、N的模平方向量,总显存占用不到广播实现的1%。GPU显存带宽是硬瓶颈,海量中间数据的读写会让计算单元全程等待数据加载,根本跑不满算力。
  • 矩阵乘法有硬件级极致优化:代码里的A @ B.t()会直接调用cuBLAS库的通用矩阵乘法(GEMM)内核,这类内核是厂商针对GPU架构做了数十年深度调优的:会自动做计算分块,把重复使用的数据放在高速共享内存、寄存器中,最大化Tensor Core利用率,内存访问做了对齐、合并优化,计算效率拉满。而广播实现用到的减法、平方、按维求和都是逐元素类操作,既没法调用Tensor Core,内核调度、内存访问的优化程度也远低于GEMM内核。
  • 计算密度差距悬殊:计算密度指单位内存访问能完成的浮点运算次数。广播实现的逐元素操作,每读入一对元素只做1~2次运算,属于典型的内存绑定操作,性能完全受限于显存带宽;而矩阵乘法中,载入一块数据可以复用多次完成大量乘加运算,属于计算绑定操作,能充分利用GPU的浮点算力。哪怕两者理论浮点运算量处于同一量级,实际运行速度也会差出几十上百倍。

识别"伪向量化"低效代码的通用方法

  • 先推演中间张量规模:写代码时顺着运算逻辑顺推每一步生成的中间张量形状,如果某一步中间张量的元素总数比输入、最终输出的规模高一个数量级及以上,基本可以判定是低效写法。比如本案例中广播实现的三维中间张量,规模比输入、输出高了d倍(d通常是几十到上千),显存开销直接爆炸。
  • 优先对齐标准BLAS算子:如果计算逻辑可以拆解为矩阵乘法、矩阵向量乘这类标准BLAS基础算子,就不要用广播+逐元素操作自行拼接。所有主流深度学习框架对BLAS级算子的优化优先级最高,能调用专用计算单元、走最优化内核,性能比自行拼接的逐元素逻辑高1~2个数量级是常态。
  • 判断操作的计算密度属性:逐元素的加、减、乘、平方、激活这类操作,计算密度极低,属于内存绑定操作,堆再多这类操作也很难跑满GPU算力;而矩阵乘法、卷积这类算子计算密度高,能充分利用GPU的计算单元。如果你的实现全靠逐元素广播操作堆出来,哪怕没有显式for循环,也大概率是性能很差的"伪向量化"代码。
  • 警惕无意义的维度扩展:如果广播逻辑需要把低维张量扩展到更高维度、把张量总元素数放大几十上百倍才能完成计算,几乎都可以通过数学公式变形,转化为不需要高维中间张量的低维算子组合。比如本案例中通过平方差展开,直接砍掉了三维中间张量的所有开销。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 16:09:09