PyTorch多进程共享内存下matmul运算慢30倍问题排查
多进程下PyTorch matmul性能暴跌的原因与解决方案
这确实是个非常反直觉的性能问题,我来帮你拆解最可能的原因和对应的解决思路:
最可能的元凶:PyTorch线程池与多进程的资源竞争
PyTorch在CPU上执行密集运算(比如matmul)时,默认会启用多线程加速(依赖MKL、OpenBLAS这类BLAS后端的线程池)。默认情况下,线程池的大小等于CPU核心数——在你的8核机器上,单进程会用8个线程来跑运算。
当你启动2个进程后,每个进程都会独立创建自己的8线程池,总线程数瞬间变成16,远超你的8个物理核心。这会导致:
- 频繁的CPU上下文切换,消耗大量额外资源
- 每个运算线程都在争抢CPU时间片,无法持续利用缓存
- BLAS后端的线程调度逻辑混乱,反而大幅降低运算效率
这完全能解释你看到的30倍性能下降——单进程时线程池高效利用核心,多进程时线程过载导致运算彻底“堵死”。
其他可能的辅助因素
- CPU缓存竞争:如果两个进程被调度到同一个物理核心的超线程上,它们会共享L1/L2缓存。
matmul是缓存敏感型运算,单进程时独占缓存命中率高,多进程时缓存被平分,大量中间数据无法命中缓存,导致运算速度暴跌。 - 共享内存的隐藏开销:虽然PyTorch的共享内存张量理论上是直接访问,但只读张量在某些BLAS实现中可能仍会触发额外的内存同步操作(比如内存栅栏),尤其是当子进程访问主进程创建的共享张量时,可能有隐藏的锁机制导致延迟。
验证与解决步骤
1. 限制每个进程的线程数(最优先尝试)
在每个子进程启动时,强制设置PyTorch的线程数为1,避免多进程的线程池冲突:
import torch import torch.multiprocessing as mp def worker(shared_weight): # 关键:限制当前进程的PyTorch线程数为1 torch.set_num_threads(1) # 也可以通过环境变量设置:os.environ["MKL_NUM_THREADS"] = "1" input_tensor = torch.randn(1, 24, 180) # 执行matmul运算 output = input_tensor.matmul(shared_weight.t()) # ... 其他逻辑 if __name__ == '__main__': weight = torch.randn(31, 180) weight.share_memory_() # 放入共享内存 p1 = mp.Process(target=worker, args=(weight,)) p2 = mp.Process(target=worker, args=(weight,)) p1.start() p2.start() p1.join() p2.join()
这一步几乎能解决90%以上的多进程CPU运算性能问题,你可以先测试这个方案是否能让matmul的耗时回到接近单进程的水平。
2. 绑定进程到独立物理核心
如果限制线程数后仍有性能问题,可以尝试将每个子进程绑定到不同的物理核心,避免缓存竞争:
import os def worker(shared_weight, core_id): torch.set_num_threads(1) # 绑定当前进程到指定核心(Linux/macOS) os.sched_setaffinity(0, {core_id}) # Windows可以用win32api实现类似功能 # ... 运算逻辑 # 启动时指定不同核心 p1 = mp.Process(target=worker, args=(weight, 0)) p2 = mp.Process(target=worker, args=(weight, 1))
3. 对比测试:放弃共享内存,每个进程独立加载权重
如果共享内存的隐藏开销确实存在,可以试试让每个子进程单独加载权重(虽然会占用更多内存,但能避开共享内存的潜在问题),对比性能差异,验证是否是共享内存导致的问题。
4. 切换BLAS后端
不同的BLAS后端在多进程下的表现差异很大:MKL的多线程优化更激进,但多进程冲突更严重;OpenBLAS相对更稳定。你可以通过环境变量切换后端:
export TORCH_BLAS=openblas
然后重新运行代码,观察性能变化。
内容的提问来源于stack exchange,提问作者Jeremy Salwen
相关产品推荐
相关产品推荐

