调试处于可中断睡眠状态的Python并行MPI程序
调试处于可中断睡眠状态的Python并行MPI程序
看起来你遇到了MPI并行程序里的经典头疼问题——明明CPU资源充足、内存也够,进程却大部分时间卡在可中断睡眠(S状态),CPU利用率上不去。结合你提到的代码结构和依赖库,我给你分享几个实用的排查思路,都是Python+MPI场景下能用的:
1. 先搞清楚进程到底在等什么(对应C程序的strace方法)
咱不用局限于C的工具,Python也有对应的办法:
- 用strace跟踪系统调用:找一个处于S状态的进程PID,执行
strace -p <PID>,盯着输出看。如果看到大量recvfrom、sendto这类网络调用,说明进程在等MPI通信;如果是wait4或者futex,可能是在等子线程/子进程(虽然你说开了单线程,但得确认依赖库没偷偷开);如果是文件I/O相关的调用(比如read/write),那可能是磁盘慢或者代码里有文件操作瓶颈。 - 用py-spy看Python层面的调用栈:这是Python专用的采样分析器,不用改代码就能attach到进程。执行
py-spy record -o profile_rank_<PID>.svg -p <PID>,它会生成一张火焰图,能直观看到进程在睡眠时卡在哪个Python函数里——比如是卡在mpi4py.MPI.Comm.Scatter,还是卡在onnxruntime.InferenceSession.run,一目了然。
2. 给MPI程序做正确的性能分析(解决cProfile的问题)
你担心cProfile和MPI混在一起会乱,其实只要让每个进程生成独立的profile文件就行:
- 直接在启动命令里用
%p(进程ID)作为文件名后缀,比如:
这样每个MPI进程会生成以自己PID命名的mpiexec -np 16 python3 -m cProfile -o prof_%p.prof myscript.pyprof_xxxx.prof文件,完全不会冲突。 - 之后分析的时候,针对不同角色的进程分开看:比如rank0的profile重点看通信相关的函数(bcast、scatter的耗时),普通rank的profile重点看计算部分(onnx、numba、keras的调用耗时)。用
pstats模块就能分析:import pstats stats = pstats.Stats('prof_xxxx.prof') stats.sort_stats('cumulative').print_stats(20) # 看累计耗时前20的函数
3. 结合你的代码结构针对性排查
你的代码流程是bcast → scatter+计算+gather ×2,结合依赖库,重点查这几个点:
- 确认所有依赖库真的是单线程运行:你说设了
OPENMPI_NUM_THREADS=1,但很多Python计算库有自己的线程配置,必须显式在代码里设置:- TensorFlow/Keras:在代码开头加
import tensorflow as tf tf.config.threading.set_intra_op_parallelism_threads(1) tf.config.threading.set_inter_op_parallelism_threads(1) - ONNX Runtime:创建会话时指定单线程
import onnxruntime as ort sess_options = ort.SessionOptions() sess_options.intra_op_num_threads = 1 sess_options.inter_op_num_threads = 1 ort_session = ort.InferenceSession('model.onnx', sess_options=sess_options, providers=['CPUExecutionProvider']) - Numba:
import numba; numba.set_num_threads(1)
如果这些库偷偷开了多线程,128个MPI进程再各自开几个线程,256核的机器会瞬间被线程占满,导致大量上下文切换,进程就会频繁进入S状态等CPU时间片——这是很多人踩过的坑!
- TensorFlow/Keras:在代码开头加
- 验证MPI通信是不是瓶颈:把计算部分(onnx/numba/keras的代码)临时替换成空循环或者简单的算术运算,再跑程序看CPU利用率。如果替换后CPU跑满了,说明问题在计算库的调用里;如果还是卡,那就是MPI通信的问题——比如rank0在生成scatter数据时太慢,导致其他rank一直在等它发数据,这时候看rank0的cProfile,重点看生成待scatter数据的代码段耗时。
- 检查rank0的负载:rank0既要做bcast、scatter的发起,有没有额外的计算任务?如果rank0的计算量太大,它会跟不上其他rank的节奏,导致所有rank都在等它,自然会进入S状态。
4. 几个快速验证的小技巧
- 用
top -H看每个MPI进程的线程数:如果某个Python进程有多个线程,说明之前的单线程配置没生效,这肯定会导致CPU争用。 - 临时加关键日志:在代码的关键节点(比如bcast后、scatter前、计算开始/结束、gather前)加一行日志:
看哪个步骤之间的时间间隔特别长,快速定位是通信卡了还是计算卡了。import time print(f"Rank {rank} reached [step name] at {time.strftime('%H:%M:%S')}", flush=True)
内容来源于stack exchange
相关产品推荐
相关产品推荐

