如何解决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资源。这些资源无法在子进程中正确复用,会导致子进程在创建新模型时陷入死锁或资源争夺,最终表现为程序挂起。
修复方法
用
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')强制使用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')子进程显式清理残留资源
若主进程必须保留模型初始化,可在子进程创建新模型前清理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
相关产品推荐
相关产品推荐

