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

多进程处理含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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 07:15:23