多进程处理含Polars DataFrame的字典时内存过高、运行缓慢问题排查
问题描述
需要对大型Polars DataFrame按ID列唯一值拆分,基于时间划分区间。已将数据存入字典(键为ID唯一值,值为对应Polars DataFrame),尝试用多进程调用函数处理各子DataFrame,但代码运行极慢且耗尽内存,怀疑多进程实现存在问题。
现有多进程代码
import multiprocessing as mp import os import time from tqdm import tqdm import polars as pl cores = mp.cpu_count() def process_identifier(i): time.sleep(0.1) return os.getpid() if __name__ == '__main__': with mp.Pool(processes=cores) as pool: # Use chunksize to balance load process_identifier_list = pool.map(process_identifier, range(cores)) process_identifier_dict = {id: idx for idx, id in enumerate(process_identifier_list)} task_list = [(process_identifier_dict, value) for key,value in df_dict.items()] print("Simulation started") results = list(tqdm(pool.imap_unordered(func_interval_divider, task_list), total=len(task_list))) df_new = pl.concat(results)
现有处理函数
import numpy as np import polars as pl def func_interval_divider(df): print(df.columns) print(df) interval_counter = 0 interval_length = 480 patient_index_list = df.get_column('patient_index').to_numpy() for i in patient_index_list: #start of the data if i == 0: df = df.with_columns(INTERVAL_ID = pl.when( pl.col('patient_index').is_between(0, interval_length)) .then(pl.lit(interval_counter)) .otherwise(pl.col('INTERVAL_ID') ) ) interval_counter += 1 #rest of the data if (i+interval_length >= np.max(patient_index_list) and df.filter(pl.col('patient_index')==i).select(pl.col("INTERVAL_ID")).item() == 0) or (i % interval_length ==0 and i!=0): interval_counter += 1 df = df.with_columns(INTERVAL_ID = pl.when( pl.col('patient_index').is_between(i, i + interval_length)) .then(pl.lit(interval_counter)) .otherwise(pl.col('INTERVAL_ID') ) ) else: pass for i in patient_index_list: #SAE interval_id if df.filter(pl.col('patient_index')==i).select(pl.col("SAE_INTERVAL")).item() == 1 and df.filter(pl.col('patient_index')==int(i-1)).select(pl.col("SAE_INTERVAL")).item() == 0: interval_counter += 1 df = df.with_columns(INTERVAL_ID = pl.when( pl.col('patient_index').is_between(i, i + interval_length), pl.col('SAE_INTERVAL') == 1) .then(pl.lit(interval_counter)) .otherwise(pl.col('INTERVAL_ID') ) ) # edge case where two SAE_intervals follow each other if df.filter(pl.col('patient_index')==i).select(pl.col("SAE_INTERVAL")).item() == 1 and df.filter(pl.col('patient_index')==i+1).select(pl.col("SAE_INTERVAL")).item() == 1 and df.filter(pl.col('patient_index')==i).select(pl.col("INTERVAL_ID")).item()> df.filter(pl.col('patient_index')==i+1).select(pl.col("INTERVAL_ID")).item(): interval_counter += 1 df = df.with_columns(INTERVAL_ID = pl.when( pl.col('patient_index').is_between(i, i + interval_length), pl.col('SAE_INTERVAL') == 1) .then(pl.lit(interval_counter)) .otherwise(pl.col('INTERVAL_ID') ) ) else: pass return df
区间划分规则
- 每x个索引,interval_id自增1,后续x行设为该值;
- 遇到新PATIENT_ID时,interval_id重置为0;
- 当SAE列值为1时,立即生成新interval_id,后续x行设为该值,SAE值为0时停止;
- 连续SAE=1超过x行时,每x行interval_id自增1,重复此规则直至SAE=0。
优化方案
1. 多进程实现问题修复
- 移除冗余操作:原代码中获取进程ID映射并传递到任务列表的逻辑完全多余,直接传递子DataFrame即可,减少内存开销和初始化时间。
- 简化任务传递:避免传递不必要的字典参数,每个任务仅传入需要处理的子DataFrame。
- 规范全局变量使用:确保
df_dict在if __name__ == '__main__'块内定义,防止多进程重复初始化导致内存泄漏。
优化后的多进程代码:
import multiprocessing as mp from tqdm import tqdm import polars as pl if __name__ == '__main__': # 在此处定义或加载df_dict # df_dict = ... cores = mp.cpu_count() with mp.Pool(processes=cores) as pool: print("Simulation started") results = list(tqdm(pool.imap_unordered(func_interval_divider, df_dict.values()), total=len(df_dict))) df_new = pl.concat(results)
2. 处理函数性能优化(核心问题)
原处理函数的最大瓶颈是循环逐行处理Polars DataFrame,完全违背了Polars的向量化设计理念,这才是运行缓慢的主要原因。以下是向量化重写方案:
import polars as pl def func_interval_divider(df): interval_length = 480 # 初始化INTERVAL_ID列 df = df.with_columns(INTERVAL_ID=pl.lit(0)) # 规则1:基础区间划分(每interval_length个索引自增) df = df.with_columns( base_interval = pl.col('patient_index') // interval_length ) # 规则3&4:SAE触发的区间划分 # 标记SAE起始点(当前为1且前一个为0) df = df.with_columns( sae_start = pl.when( (pl.col('SAE_INTERVAL') == 1) & (pl.col('SAE_INTERVAL').shift(1) == 0) ).then(pl.lit(1)).otherwise(0) ) # 计算SAE区间的分组ID:每个SAE起始点对应一个新分组 df = df.with_columns( sae_group = pl.col('sae_start').cumsum() ) # 对每个SAE分组,计算内部的区间(每interval_length个索引自增) df = df.with_columns( sae_interval = pl.when(pl.col('SAE_INTERVAL') == 1) .then((pl.col('patient_index') - pl.col('patient_index').filter(pl.col('sae_start') == 1).first().over('sae_group')) // interval_length) .otherwise(0) ) # 合并基础区间和SAE区间:SAE区间优先级更高 df = df.with_columns( INTERVAL_ID = pl.when(pl.col('SAE_INTERVAL') == 1) .then(pl.col('sae_group') + pl.col('sae_interval')) .otherwise(pl.col('base_interval')) ) # 规则2:每个PATIENT_ID的interval_id从0开始(当前子df为单PATIENT_ID数据) df = df.with_columns( INTERVAL_ID = pl.col('INTERVAL_ID') - pl.col('INTERVAL_ID').min() ) # 清理临时列 df = df.drop(['base_interval', 'sae_start', 'sae_group', 'sae_interval']) return df
优化说明
- 完全向量化操作:利用Polars的窗口函数、累积求和等方法实现批量计算,避免所有循环和逐行
filter,性能提升数个数量级。 - 内存高效:无需将列转为numpy数组,减少内存拷贝;链式调用
with_columns比循环创建新DataFrame更节省内存。 - 逻辑清晰:按规则分步实现,便于维护和调试。
内容的提问来源于stack exchange,提问作者Kloevedal97
相关产品推荐
相关产品推荐

