如何用PyTorch TorchMeta创建元学习分布式数据加载器及死锁排查
死锁产生原因
- TorchMeta默认的数据集采样逻辑没有适配DDP的分布式分片规则,不同进程会重复拉取全量任务批次,使用
num_workers>0的多进程数据加载时,各进程间的采样器同步逻辑冲突,触发IO阻塞死锁 - DDP默认要求所有进程执行完全一致的集体通信操作,TorchMeta内置的
MetaDataset部分实现会在主进程预加载元任务缓存,子进程不会触发对应缓存同步逻辑,导致不同进程执行时序不一致,卡在barrier等待阶段 - 若使用TorchMeta默认的
BatchMetaDataLoader,其内置的任务采样器没有做分布式分片,不同进程会争夺同一个数据读取句柄,多进程场景下出现锁竞争 - PyTorch默认多进程启动方式为
fork,TorchMeta的元数据集加载逻辑包含大量序列化操作,和DDP的fork启动模式冲突,容易触发进程间死锁
修复方案
1. 自定义分布式元任务采样器
替换TorchMeta默认的采样器,适配DDP的分片逻辑,核心实现如下:
import torch import torch.distributed as dist from torchmeta.utils.data import BatchSampler class DistributedMetaBatchSampler(BatchSampler): def __init__(self, sampler, batch_size, drop_last, rank=None, num_replicas=None): super().__init__(sampler, batch_size, drop_last) self.rank = rank if rank is not None else dist.get_rank() self.num_replicas = num_replicas if num_replicas is not None else dist.get_world_size() # 对总任务数做分片,每个进程仅拉取自身负责的任务段 self.total_size = len(self.sampler) self.num_samples = self.total_size // self.num_replicas if self.total_size % self.num_replicas != 0 and not drop_last: self.num_samples += 1 def __iter__(self): # 每个epoch重置种子保证所有进程采样逻辑对齐 seed = int(torch.empty((), dtype=torch.int64).random_().item()) self.sampler.set_epoch(seed) batches = list(super().__iter__()) # 仅返回当前rank对应的分片批次 offset = self.rank * self.num_samples return iter(batches[offset: offset + self.num_samples])
初始化数据加载器时传入自定义采样器即可。
2. 调整多进程启动与数据加载参数
- 启动DDP进程前显式设置多进程启动模式为
spawn,避免fork模式下的序列化冲突:
if __name__ == "__main__": torch.multiprocessing.set_start_method('spawn') # 后续DDP初始化与训练逻辑
- 初始化
MetaDataLoader时指定pin_memory=False,DDP多进程下pin_memory的内存拷贝逻辑容易和TorchMeta的缓存逻辑冲突导致死锁 - 先将
num_workers设为0验证是否为多进程加载问题,确认问题后再逐步调整worker数量,同时设置worker_init_fn保证每个worker的随机种子和进程rank对齐
3. 禁用TorchMeta全局缓存逻辑
初始化MetaDataset时指定dataset_cache=False,关闭全局缓存逻辑,避免不同进程的缓存读写冲突,若需要缓存可自行实现每个rank独立的本地缓存逻辑。
内容的提问来源于stack exchange,提问作者Charlie Parker
相关产品推荐
相关产品推荐

