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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 20:15:14