PyTorch DataLoader设置num_workers大于1时出现SSLError如何解决?
问题根因
PyTorch DataLoader在Linux/macOS环境下默认使用fork模式创建子进程,你在主进程中初始化的全局requests.Session对象会被所有worker进程继承,多个进程同时操作同一个SSL连接池的套接字,会导致SSL连接状态错乱,从而抛出DECRYPTION_FAILED_OR_BAD_RECORD_MAC错误。单worker场景下只有一个进程使用session,不存在资源竞争,因此运行正常。
解决方案
你可以选择以下任意一种方案解决问题:
- 方案一(推荐):给每个worker进程初始化专属的session,避免资源竞争
你可以通过DataLoader的worker_init_fn钩子统一实现,代码修改如下:
首先调整Dataset定义:
class MyDataset(Dataset): def __init__(self, obj_ids = []): super().__init__() self.obj_ids = obj_ids # 不在主进程初始化session,留给worker进程自己创建 self.session = None def __len__(self): return len(self.obj_ids) def __getitem__(self, idx): if torch.is_tensor(idx): idx = idx.tolist() result = self.session.get('/api/url/{}'.format(idx)) # 后续处理逻辑...
新增worker初始化函数,初始化DataLoader时传入:
def worker_init_fn(worker_id): import requests # 获取当前worker对应的dataset实例 worker_info = torch.utils.data.get_worker_info() dataset = worker_info.dataset # 给当前worker的dataset实例创建专属session dataset.session = requests.Session() data_loader = torch.utils.data.DataLoader( dataset, batch_size=2, shuffle=True, num_workers=4, collate_fn=utils.collate_fn, worker_init_fn=worker_init_fn)
这种方案可以复用每个worker的session连接,性能最优。
- 方案二:直接在
__getitem__中使用无session的请求,不复用连接
如果你不想修改DataLoader配置,可以直接把session.get改成requests.get,每次请求都会新建独立的连接,不会出现多进程冲突,缺点是请求量较大时重复建立TCP连接会拉低数据加载效率:
# 修改__getitem__中的请求逻辑 import requests result = requests.get('/api/url/{}'.format(idx))
- 方案三:修改DataLoader的进程创建模式为spawn
你可以在初始化DataLoader时指定multiprocessing_context='spawn',spawn模式下子进程不会继承主进程的资源,自然不会出现session共享冲突,缺点是进程启动速度比fork模式慢,适合worker数量较少的场景:
data_loader = torch.utils.data.DataLoader( dataset, batch_size=2, shuffle=True, num_workers=4, collate_fn=utils.collate_fn, multiprocessing_context='spawn')
注意事项
- 不要在Dataset的
__init__方法中初始化requests session,主进程创建的session一定会被子进程继承导致冲突 - 如果调用的API有请求频率限制,需要额外给每个worker添加限流逻辑,避免多worker并发请求触发接口限流
内容的提问来源于stack exchange,提问作者Pablo Estrada
相关产品推荐
相关产品推荐

