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

如何解决stable_baselines3多进程初始化模型时的代码挂起问题

尝试在各自进程中初始化多个stable_baselines3模型

以下代码使用Pool初始化模型时会挂起,若注释掉initialize_model(1)行,代码可正常运行完成。

在多进程前运行initialize_model(1)会引发什么问题?是否有修复方法?

(注:已知stable_baselines3自带并行化功能,但因业务特殊性无法使用)

from multiprocessing import Pool
# pathos has same issue
# from pathos.multiprocessing import Pool
from stable_baselines3 import PPO
from stable_baselines3.ppo import MlpPolicy


def initialize_model(t):
    model = PPO(policy=MlpPolicy, env='CartPole-v1', batch_size=128, n_steps=128)
    return 1


# TODO: with the following line uncommented, it does not work
initialize_model(1)

with Pool(processes=2) as pool:
    test = pool.map(initialize_model, [None for _ in range(2)])
print('done')

问题原因

主进程提前初始化PPO模型时,会触发PyTorch的CUDA上下文(即便没显式指定GPU,底层部分逻辑仍会完成初始化)或共享资源的创建。当multiprocessing.Pool通过fork机制创建子进程时,会完整复制主进程的内存空间,包括已初始化的PyTorch/CUDA资源。这些资源无法在子进程中正确复用,会导致子进程在创建新模型时陷入死锁或资源争夺,最终表现为程序挂起。

修复方法

  1. 用if __name__ == '__main__':包裹主进程逻辑
    将主进程的模型初始化和多进程代码放在该判断块内,避免子进程在启动时重复执行初始化操作,阻断无效的资源复制:

    from multiprocessing import Pool
    from stable_baselines3 import PPO
    from stable_baselines3.ppo import MlpPolicy
    
    
    def initialize_model(t):
        model = PPO(policy=MlpPolicy, env='CartPole-v1', batch_size=128, n_steps=128)
        return 1
    
    
    if __name__ == '__main__':
        initialize_model(1)
    
        with Pool(processes=2) as pool:
            test = pool.map(initialize_model, [None for _ in range(2)])
        print('done')
    
  2. 强制使用spawn启动模式
    若必须在主进程提前初始化模型,可强制Pool使用spawn模式创建子进程。该模式会启动全新的Python进程,不会复制主进程内存空间,从根源避免共享资源冲突:

    from multiprocessing import Pool, get_context
    from stable_baselines3 import PPO
    from stable_baselines3.ppo import MlpPolicy
    
    
    def initialize_model(t):
        model = PPO(policy=MlpPolicy, env='CartPole-v1', batch_size=128, n_steps=128)
        return 1
    
    
    if __name__ == '__main__':
        initialize_model(1)
        with get_context('spawn').Pool(processes=2) as pool:
            test = pool.map(initialize_model, [None for _ in range(2)])
        print('done')
    
  3. 子进程显式清理残留资源
    若主进程必须保留模型初始化,可在子进程创建新模型前清理PyTorch残留资源(适配性稍差,依赖PyTorch版本):

    import torch
    from multiprocessing import Pool
    from stable_baselines3 import PPO
    from stable_baselines3.ppo import MlpPolicy
    
    
    def initialize_model(t):
        torch.cuda.empty_cache()
        if torch.cuda.is_available():
            torch.cuda.device_reset()
        model = PPO(policy=MlpPolicy, env='CartPole-v1', batch_size=128, n_steps=128)
        return 1
    
    
    initialize_model(1)
    with Pool(processes=2) as pool:
        test = pool.map(initialize_model, [None for _ in range(2)])
    print('done')
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 01:52:39