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

如何在Slurm集群上实现Python任务中间阶段的并行计算?

如何在Slurm集群上实现Python任务中间阶段的并行计算?

首先得搞清楚你遇到的Slurm警告到底是怎么回事:你请求了4个节点,但Slurm发现你的任务只启动了1个主Python进程,而这个进程里的多进程(pathos的Pool)只能在单个节点内跑,跨节点的话子进程没法通信,所以Slurm自动把节点数改成1了——这其实是Slurm在帮你避免资源浪费。

你的需求是“任务中间某一步才需要并行”,这种场景其实不用多节点,单节点内的多进程就足够(尤其是你说n大概200,这个规模单节点的CPU完全能搞定),下面给你两种适配Slurm的解决方案:

方案一:单节点内多进程(最适合你的当前场景)

1. 调整Slurm提交脚本

你只需要申请单个节点,同时申请足够的CPU核心数(和你代码里Pool的进程数对应),比如下面的脚本:

#!/bin/bash
#SBATCH --job-name=wave_mode_calc
#SBATCH --nodes=1          # 只需要单节点
#SBATCH --cpus-per-task=4  # 给主进程分配4个核心,对应Pool的进程数
#SBATCH --time=00:20:00    # 根据你的计算时间调整
#SBATCH --output=job_%j.out

python your_script.py

2. 优化Python代码(动态适配Slurm资源)

不要硬编码Pool的进程数,而是读取Slurm的环境变量SLURM_CPUS_PER_TASK,这样脚本申请多少核心,代码就用多少,更灵活:

class Foo:
    def __init__(self, n):
        self.n = n
        self.nList = list(range(n))  # 关于你的Bonus问题:这个写法已经很优雅了,如果用numpy可以写成np.arange(n).tolist(),但原生Python里list(range(n))就是最简洁可读的
    
    def cubicRoot(self, x):
        # 这里替换成你的波方程模式计算
        return x**(1/3)
    
    def cubicRootParallel(self):
        import os
        from pathos.multiprocessing import ProcessingPool as Pool
        # 从Slurm环境变量拿核心数,默认4(本地测试用)
        num_workers = int(os.getenv('SLURM_CPUS_PER_TASK', 4))
        p = Pool(num_workers)
        
        def _cubicRoot(x):
            return self.cubicRoot(x)
  
        self.cubicRootList = p.map(_cubicRoot, self.nList)

foo = Foo(200)
foo.cubicRootParallel()  # 你原代码漏了括号!这里要调用方法才能执行
print(foo.cubicRootList)

👉 注意:你原代码里漏写了foo.cubicRootParallel()的括号,这会导致方法没被执行,cubicRootList根本不会被生成,这个小bug得先改掉!

方案二:跨节点并行(未来大规模计算用)

如果以后你的n变得特别大,单节点不够用了,那就要用分布式计算框架,比如mpi4py,结合Slurm的多节点设置。这种方式是让Slurm在每个节点启动多个MPI进程,每个进程处理一部分任务,最后汇总结果:

1. Python代码示例(用mpi4py)

from mpi4py import MPI
import numpy as np

class Foo:
    def __init__(self, n):
        self.n = n
        self.nList = list(range(n))

    def cubicRoot(self, x):
        # 替换成你的波方程模式计算
        return x**(1/3)
    
    def cubicRootParallel(self):
        comm = MPI.COMM_WORLD
        rank = comm.Get_rank()  # 当前进程的编号
        total_processes = comm.Get_size()  # 总进程数

        # 把任务平均分给每个进程
        local_tasks = np.array_split(self.nList, total_processes)[rank]
        # 每个进程计算自己的任务
        local_results = [self.cubicRoot(x) for x in local_tasks]

        # 把所有进程的结果收集到主进程(rank=0)
        all_results = comm.gather(local_results, root=0)

        if rank == 0:
            # 合并所有结果
            self.cubicRootList = [item for sublist in all_results for item in sublist]

if __name__ == "__main__":
    foo = Foo(200)
    foo.cubicRootParallel()
    # 只有主进程打印结果
    if MPI.COMM_WORLD.Get_rank() == 0:
        print(foo.cubicRootList)

2. 对应的Slurm脚本

#!/bin/bash
#SBATCH --job-name=large_wave_calc
#SBATCH --nodes=2          # 申请2个节点
#SBATCH --ntasks-per-node=8  # 每个节点启动8个MPI进程,总共16个
#SBATCH --time=00:30:00
#SBATCH --output=job_%j.out

# 用srun启动MPI进程
srun python your_script.py

额外小建议

对于你的波方程计算,如果每个模式的计算是数组操作,优先用numpy的向量化运算,比如直接np.array(self.nList) ** (1/3),这种方式比多进程更快——因为numpy是C底层实现,没有多进程的启动和通信开销,尤其是n=200这种规模,可能单线程向量化计算比多进程还高效。

备注:内容来源于stack exchange,提问作者dolefeast

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 11:02:58