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

Python多进程池中使用SQLAlchemy遇Pickle序列化错误求助

问题原因分析

这个Pickle错误的核心原因有三点:

  1. SQLAlchemy引擎无法被序列化:SQLAlchemy的Engine对象内部包含create_engine生成的局部connect函数,pickle无法序列化局部函数。当子进程抛出异常时,异常对象可能间接引用了Engine,导致进程池传递异常结果时触发序列化失败。
  2. 跨进程共享引擎不符合规范:数据库连接是进程绑定的,父进程创建的引擎传递给子进程后,子进程无法复用连接,反而会引发资源冲突和序列化问题,你之前尝试的共享引擎做法本身就不成立。
  3. init函数表检查逻辑错误:原代码中inspect(engine).has_table(engine, 'Company')写法错误,has_table的第一个参数应该是表名而非引擎,这会导致表存在性检查失效,触发额外异常,进一步引发序列化问题。
解决方案

1. 修复init函数的表检查逻辑

修正表存在性检查的调用方式,确保逻辑正确:

from sqlalchemy import inspect

def init(db_url: str) -> sqlalchemy.engine.Engine:
    engine = create_engine(db_url, poolclass=sqlalchemy.pool.QueuePool)
    inspector = inspect(engine)
    # 修正has_table参数,传入表名而非引擎
    if not inspector.has_table('Company'):
        Base.metadata.tables['Company'].create(engine)
    if not inspector.has_table('FormerName'):
        Base.metadata.tables['FormerName'].create(engine)
    if not inspector.has_table('Address'):
        Base.metadata.tables['Address'].create(engine)
    if not inspector.has_table('Ticker'):
        Base.metadata.tables['Ticker'].create(engine)
    if not inspector.has_table('ExchangeCompany'):
        Base.metadata.tables['ExchangeCompany'].create(engine)
    if not inspector.has_table('Filing'):
        Base.metadata.tables['Filing'].create(engine)
    return engine

2. 子进程独立创建引擎,隔离异常序列化风险

在load_submission内部捕获所有异常,确保异常信息不包含Engine或Session这类无法序列化的对象:

def load_submission(submission_files: list[str], db_url: str) -> None:
    try:
        submission = extract_submission(submission_files)
        c = submission.get('cik')
        if not c:
            raise KeyError(f'No CIK found in submission')
        c = c.zfill(10)  

        log.debug(f'Process {os.getpid()} is loading submission for CIK: {c}')
        
        # 子进程独立创建引擎
        engine = init(db_url)
        log.debug(f'Engine initialized for CIK: {c}')

        with engine.connect() as conn:
            with conn.begin():
                # 处理Company数据
                company = db.Company(
                    cik = c,
                    name = submission.get('name'),
                    sic = submission.get('sic'),
                    entityType = submission.get('entityType'),
                    insiderTransactionForOwnerExists = submission.get('insiderTransactionForOwnerExists'),
                    insiderTransactionForIssuerExists = submission.get('insiderTransactionForIssuerExists'),
                    ein = submission.get('ein'),
                    description = submission.get('description'),
                    website = submission.get('website'),
                    investorWebsite = submission.get('investorWebsite'),
                    category = submission.get('category'),
                    fiscalYearEnd = submission.get('fiscalYearEnd'),
                    stateOfIncorporation = submission.get('stateOfIncorporation'),
                    phone = submission.get('phone'),
                    flags = submission.get('flags')
                )
                conn.merge(company)

                # 处理Ticker数据
                for ticker in submission['tickers']:
                    t = db.Ticker(ticker=ticker, cik=c)
                    conn.merge(t)

                # 处理ExchangeCompany数据
                for exchange in submission['exchanges']:
                    ec = db.ExchangeCompany(exchange=exchange, cik=c)
                    conn.merge(ec)

                # 处理Address数据
                addresses = submission['addresses']
                for description, address in addresses.items():
                    a = db.Address(
                        cik = c,
                        description = description,
                        street1 = address.get('street1'),
                        street2 = address.get('street2'),
                        city = address.get('city'),
                        stateOrCountry = address.get('stateOrCountry'),
                        zipCode = address.get('zipCode')
                    )
                    conn.merge(a)

                # 处理FormerName数据
                former_names = submission['formerNames']
                for former_name in former_names:
                    f = db.FormerName(
                        formerName = former_name.get('name'),
                        _from = datetime.fromisoformat(former_name.get('from')) if former_name.get('from') else None,
                        to = datetime.fromisoformat(former_name.get('to')) if former_name.get('to') else None,
                        cik = c
                    )
                    conn.merge(f)

                # 处理Filing数据
                df = pd.DataFrame(submission['filings']['recent'])
                if not df.empty:
                    df['url'] = df.apply(lambda row: f'{ARCHIVES_URL}/{str(int(c))}/{row["accessionNumber"].replace("-", "")}/{row["accessionNumber"]}.txt', axis=1)
                    df['cik'] = c

                    # 日期转换
                    date_columns = ['filingDate', 'reportDate', 'acceptanceDateTime']
                    for col in date_columns:
                        df[col] = pd.to_datetime(df[col], errors='coerce')

                    dtypes = {
                        'cik': sqltypes.VARCHAR(),
                        'accessionNumber': sqltypes.VARCHAR(),
                        'filingDate': sqltypes.Date(),
                        'reportDate': sqltypes.Date(),
                        'acceptanceDateTime': sqltypes.DateTime(),
                        'act': sqltypes.VARCHAR(),
                        'fileNumber': sqltypes.VARCHAR(),
                        'filmNumber': sqltypes.VARCHAR(),
                        'items': sqltypes.VARCHAR(),
                        'size': sqltypes.INTEGER(),
                        'isXBRL': sqltypes.Boolean,
                        'isInlineXBRL': sqltypes.Boolean,
                        'primaryDocument': sqltypes.VARCHAR(),
                        'primaryDocumentDescription': sqltypes.VARCHAR()
                    }
                    df.to_sql('Filing', conn, if_exists='append', index=False, dtype=dtypes)
        
        engine.dispose()
        log.debug(f'Process {os.getpid()} finished loading CIK: {c}')
    except Exception as e:
        # 捕获异常并记录,避免异常携带无法序列化对象
        log.error(f'Failed to load submission for files {submission_files}: {str(e)}')
        # 抛出不含引擎引用的简化异常
        raise RuntimeError(f'Failed to load submission: {str(e)}') from None

3. 优化进程池初始化(可选,降低引擎创建开销)

如果每个任务创建引擎的开销过大,可以用进程池的initializer让每个子进程只初始化一次引擎:

# 子进程全局变量存储引擎
worker_engine = None

def init_worker(db_url):
    global worker_engine
    worker_engine = init(db_url)

def load_submission(submission_files: list[str]) -> None:
    try:
        submission = extract_submission(submission_files)
        c = submission.get('cik')
        if not c:
            raise KeyError(f'No CIK found in submission')
        c = c.zfill(10)  

        log.debug(f'Process {os.getpid()} is loading submission for CIK: {c}')
        
        # 使用子进程全局引擎
        engine = worker_engine
        # 后续数据处理逻辑同上...
        
    except Exception as e:
        log.error(f'Failed to load submission for files {submission_files}: {str(e)}')
        raise RuntimeError(f'Failed to load submission: {str(e)}') from None

# 主进程创建池时指定初始化函数
def load_submissions(submissions_path:str, db_url:str) -> None:
    submissions = {}
    for file in os.listdir(submissions_path):
        file_path = os.path.abspath(os.path.join(submissions_path, file))
        file = os.path.basename(file)
        cik = numbers_from_filename(file)
        cik = cik[:10] if cik else None
        
        if cik:
            submissions.setdefault(cik, []).append(file_path)
        else:
            log.info(f'Could not find a CIK in the filename: {file_path}')
    log.debug(f'Submissions loaded')

    tasks = [files for files in submissions.values()]
    try:
        # 根据数据库连接数限制进程数,避免连接耗尽
        num_processes = min(len(submissions), multiprocessing.cpu_count(), 10)
        with multiprocessing.get_context("spawn").Pool(
            processes=num_processes,
            initializer=init_worker,
            initargs=(db_url,)
        ) as pool:
            log.debug(f'Loading submissions with multiprocessing.Pool()')
            pool.map(load_submission, tasks)
    except Exception as e:
        log.exception(e)

4. 额外性能优化建议

  • 控制进程数量:不要直接用CPU核心数,参考PostgreSQL的max_connections配置,建议设置为min(len(submissions), max_connections // 2),避免数据库连接池耗尽。
  • 使用COPY协议加速写入:对于Filing表的批量写入,改用psycopg2的copy_from或SQLAlchemy的copy_expert方法,比to_sql的插入效率提升数倍。
  • 批量分组任务:如果单个CIK的文件数量过少,可以合并多个CIK为一个任务,减少进程调度开销。

内容的提问来源于stack exchange,提问作者Michael

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 17:25:55