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

PyTorch多GPU预测代码在spawn()处冻结的问题排查求助

PyTorch多GPU预测代码在spawn()处冻结的问题排查求助

大家好,我最近在写PyTorch多GPU预测的代码,遇到了一个棘手的问题,想请各位帮忙分析一下:

我用torch.multiprocessing.spawn()启动了4个进程做多GPU预测,任务函数里没有写实际的预测逻辑,只是把处理后的数据放到一个公共的Queue里用来收集结果。现在所有进程都打印出了{rank} destroyed的消息,说明每个进程都成功销毁了进程组,但spawn()之后的Spawn end消息却永远不会出现——这意味着torch.multiprocessing.spawn()根本没有结束。我在任务函数里加了try-catch块,也没有捕获到任何异常,实在找不到问题所在了。

以下是我的完整代码:

from typing import List

import torch
import torch.multiprocessing
import torch.distributed
import os
from torch.utils.data.dataloader import DataLoader
from torch.utils.data.distributed import DistributedSampler

world_size = 4
batch_size = 16
data = [i for i in range(10000)]

class DistributedDataset(torch.utils.data.Dataset):
def __init__(self, list:List):
    super(DistributedDataset, self).__init__()
    self.data = list

def __getitem__(self, index):
    return self.data[index]

def __len__(self):
    return len(self.data)

dataset = DistributedDataset(list=data)

def task(rank:int, result_queue:torch.multiprocessing.Queue):
try:
    torch.distributed.init_process_group('nccl', rank=rank, world_size=world_size)
    torch.cuda.set_device(device=rank)
    
    sampler = DistributedSampler(dataset, num_replicas=world_size, rank=rank)
    dataloader = DataLoader(dataset,
                            batch_size=batch_size,
                            pin_memory=True,
                            sampler=sampler,
                            shuffle=False)
    
    process_result = []
    for batch in dataloader:
        process_result.extend(batch)
    
    result_queue.put(process_result)
    
    torch.distributed.barrier()
    torch.distributed.destroy_process_group()
    print(f'{rank} destroyed')
except:
    print('error')

if __name__ == '__main__':
    os.environ['MASTER_ADDR'] = 'localhost'
    os.environ['MASTER_PORT'] = '12356'
    
    torch.multiprocessing.set_start_method('spawn', force=True)
    result_queue = torch.multiprocessing.Queue()
    
    torch.multiprocessing.spawn(task, args=(result_queue,), nprocs=world_size)
    print(f'Spawn end')
    
    results = []
    for _ in range(world_size):
        results.append(result_queue.get())

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.21 07:13:00