使用torch.multiprocessing时遭遇死锁问题求助
PyTorch多进程初始化死锁问题的解决方案
问题重现
运行以下代码时出现死锁:
import torch.multiprocessing as mp import torch from torch import nn import numpy as np class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.vars = nn.ParameterList() print("Net init 0") weight = nn.Parameter(torch.nn.init.orthogonal_(torch.zeros([64, 64]), 1.4)) print("Net init 1") bias = nn.Parameter(torch.nn.init.constant_(torch.zeros(64), 1.4)) print("Net init 2") self.vars.extend([weight, bias]) print("Net init 3") def f(): net = Net() if __name__ == "__main__": agent = Net() processes = [mp.Process(target=f) for _ in range(2)] for p in processes: p.start() for p in processes: p.join()
执行输出显示子进程卡在Net init 0后无法继续:
# python3 test.py Net init 0 Net init 1 Net init 2 Net init 3 Net init 0 Net init 0
但注释主进程的agent = Net()语句,或移除torch.nn.init.orthogonal_调用时,代码可正常运行。
原因分析
死锁源于PyTorch全局随机数生成器(RNG)的锁机制:
torch.nn.init.orthogonal_内部会调用全局RNG,此时会获取一把全局锁;- 主进程初始化
agent时,调用orthogonal_后可能未完全释放该锁; - Python的
multiprocessing基于fork机制,子进程会完整复制父进程的内存状态,包括未释放的锁; - 子进程启动后调用
orthogonal_时,尝试获取同一锁,因父进程已占用(或复制的锁状态异常)而阻塞,最终死锁。
解决方案
方案1:主进程初始化后重置RNG状态
在主进程创建完agent后,调用torch.random._fork_rng()重置RNG状态,避免子进程继承异常的锁状态:
if __name__ == "__main__": agent = Net() # 重置RNG,清除可能遗留的锁状态 torch.random._fork_rng() processes = [mp.Process(target=f) for _ in range(2)] for p in processes: p.start() for p in processes: p.join()
方案2:调整主进程初始化时机
将主进程的agent初始化移到子进程启动之后,避免fork时复制带锁的RNG状态:
if __name__ == "__main__": processes = [mp.Process(target=f) for _ in range(2)] for p in processes: p.start() # 子进程启动后再初始化主进程的Net agent = Net() for p in processes: p.join()
方案3:使用局部RNG隔离初始化逻辑
在__init__中用局部RNG上下文执行orthogonal_,避免占用全局锁:
class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.vars = nn.ParameterList() print("Net init 0") # 先创建空参数,再用局部RNG初始化 weight = nn.Parameter(torch.empty([64, 64])) with torch.random.fork_rng(): torch.nn.init.orthogonal_(weight, 1.4) print("Net init 1") bias = nn.Parameter(torch.nn.init.constant_(torch.zeros(64), 1.4)) print("Net init 2") self.vars.extend([weight, bias]) print("Net init 3")
内容的提问来源于stack exchange,提问作者he xiangdong
相关产品推荐
相关产品推荐

