Python多进程池中使用SQLAlchemy遇Pickle序列化错误求助
问题原因分析
这个Pickle错误的核心原因有三点:
- SQLAlchemy引擎无法被序列化:SQLAlchemy的
Engine对象内部包含create_engine生成的局部connect函数,pickle无法序列化局部函数。当子进程抛出异常时,异常对象可能间接引用了Engine,导致进程池传递异常结果时触发序列化失败。 - 跨进程共享引擎不符合规范:数据库连接是进程绑定的,父进程创建的引擎传递给子进程后,子进程无法复用连接,反而会引发资源冲突和序列化问题,你之前尝试的共享引擎做法本身就不成立。
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
相关产品推荐
相关产品推荐

