使用keras.utils.Sequence多进程与MongoDB时,何时建立数据库连接?
解决Keras多进程训练时MongoDB连接的问题
这个问题我之前踩过坑!核心原因是多进程模式下,父进程的MongoDB连接对象无法被安全继承到子进程中——MongoDB的连接基于TCP协议,进程fork后子进程会复制父进程的文件描述符,但这些连接在子进程里是不可用的,而且连接对象本身不支持序列化,直接在__init__里创建连接必然会触发异常。
下面是我验证过的可行解决方案:
1. 延迟创建连接,让每个进程独立初始化
不要在Sequence的__init__方法里创建MongoDB连接,而是把连接逻辑放到__getitem__(批次数据获取时)或者专门的初始化方法里,让每个工作进程自己创建专属连接。这样每个进程的连接完全独立,不会出现跨进程的冲突。
示例代码如下:
from keras.utils import Sequence from pymongo import MongoClient import numpy as np class MongoDataSequence(Sequence): def __init__(self, db_name, collection_name, batch_size): self.db_name = db_name self.collection_name = collection_name self.batch_size = batch_size # 只保存连接参数,不提前创建连接 self.client = None self.collection = None # 父进程提前获取总样本数(避免每个子进程重复查询) temp_client = MongoClient() self.total_samples = temp_client[db_name][collection_name].count_documents({}) temp_client.close() def _get_collection(self): # 每个进程第一次调用时,创建自己的连接 if self.client is None: self.client = MongoClient() # 这里可以传入你的MongoDB地址、认证信息等 self.collection = self.client[self.db_name][self.collection_name] return self.collection def __len__(self): return (self.total_samples + self.batch_size - 1) // self.batch_size def __getitem__(self, idx): collection = self._get_collection() # 计算当前批次的起止位置 start = idx * self.batch_size end = start + self.batch_size # 从MongoDB拉取批次数据 cursor = collection.find().skip(start).limit(self.batch_size) # 转换成模型需要的输入格式(这里根据你的数据结构调整) x, y = [], [] for doc in cursor: x.append(doc['features']) y.append(doc['label']) return np.array(x), np.array(y) def on_epoch_end(self): # 可选:每个epoch结束后关闭连接,避免长期占用资源 if self.client is not None: self.client.close() self.client = None self.collection = None
2. 优化:用进程本地存储复用连接
如果觉得每次__getitem__都检查连接有点繁琐,也可以用multiprocessing.local()来存储每个进程的连接,确保每个进程只有一个活跃连接:
from multiprocessing import local class MongoDataSequence(Sequence): def __init__(self, db_name, collection_name, batch_size): self.db_name = db_name self.collection_name = collection_name self.batch_size = batch_size self.local = local() # 进程本地存储,每个进程有独立的副本 # 父进程预计算总样本数 temp_client = MongoClient() self.total_samples = temp_client[db_name][collection_name].count_documents({}) temp_client.close() def _get_collection(self): if not hasattr(self.local, 'collection'): self.local.client = MongoClient() self.local.collection = self.local.client[self.db_name][self.collection_name] return self.local.collection # 其余__len__、__getitem__、on_epoch_end方法和上面一致
关键注意事项
- 绝对不要在父进程创建连接后传递给子进程:MongoDB连接是进程绑定的,跨进程复用会直接导致连接失效、抛出异常。
- 控制连接数量:如果设置了
workers=N,就会有N个工作进程,每个进程创建一个连接,确保MongoDB的maxConnections配置足够(默认值100足够大多数场景)。 - 按需清理连接:如果你的训练周期很长,可以在
on_epoch_end里关闭连接,避免连接资源被长期占用。
按照这个方式修改后,再开启use_multiprocessing=True应该就能正常运行了!
内容的提问来源于stack exchange,提问作者wl2776
相关产品推荐
相关产品推荐

