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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 06:20:16