如何高效将带过滤的Pandas透视表转换为Polars实现
Pandas转Polars:实现等效的透视表过滤逻辑
需求背景
需要将Pandas中针对透视表的过滤逻辑转换为Polars实现,原始数据集规模超500万行,已用Polars生成透视表,需完成过滤得到与Pandas中df_final一致的结果。
原始Pandas核心逻辑
先通过透视表聚合,再通过两个条件过滤:
- 按
col4分组,对每行所有列的值降序排名,保留至少有一列排名为1的行 - 保留行中至少有一个值大于100的行
原始Pandas代码:
import pandas as pd df = pd.DataFrame({ 'col1': [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19], 'col2': ['test1', 'test1', 'test1', 'test1', 'test2', 'test2', 'test2', 'test2', 'test3', 'test3', 'test3', 'test3', 'test4', 'test5', 'test1', 'test1', 'test1', 'test3', 'test4'], 'col3': ['t1', 't1', 't1', 't1', 't1', 't1', 't1', 't1', 't1', 't1', 't1', 't1', 't1', 't1', 't1', 't1','tl','tl','tl'], 'col4': ['input1', 'input2', 'input3', 'input4', 'input1', 'input2', 'input3', 'input4', 'input1', 'input2', 'input3', 'input5', 'input2', 'input6', 'input1', 'input1', 'input2', 'input2', 'input2'], 'col5': ['result1', 'result2', 'result3', 'result4', 'result1', 'result2', 'result3', 'result4', 'result1', 'result2', 'result3', 'result4', 'result2', 'result1', 'result2', 'result6', 'result1', 'result1', 'result1'], 'col6': [10, 20, 30, 40, 10, 20, 30, 40, 10, 20, 30, 50, 20, 100, 10, 10, 20, 20, 20], 'col7': [100.2, 101.2, 102.3, 101.4, 100.0, 103.0, 104.0, 105.0, 102.0, 87.0, 107.0, 110.2, 120.0, 88.0, 106.2, 101.1, 100, 90.2, 110] }) # 生成透视表 p_df = df.pivot_table(values='col7', index=['col4', 'col5', 'col6'], columns=['col2'], aggfunc='max') # 双条件过滤 df_final = p_df[((p_df.groupby(level=0).rank(ascending=False) == 1.).any(axis=1))&(p_df>100).any(axis=1)] print(df_final)
Polars等效实现
注意Polars的pivot默认聚合函数是mean,需要显式指定aggfunc=pl.max;同时利用Polars的列操作和窗口函数实现分组排名逻辑,处理大数据时建议使用Lazy模式提升效率:
import polars as pl # 读取数据(直接用Pandas生成的示例数据,实际可直接从源读取) df_p = pl.from_pandas(df) # 1. 生成透视表,指定max聚合 p_df = df_p.pivot( on='col2', index=['col4', 'col5', 'col6'], values='col7', aggregate_function=pl.max ) # 2. 定义过滤条件 # 获取所有透视后的列(即test1/test2/test3/test4/test5这些列) pivot_cols = [col for col in p_df.columns if col not in ['col4', 'col5', 'col6']] # 条件1:按col4分组,每行各列值降序排名,至少有一列排名为1 cond1 = pl.any( pl.col(pivot_cols).rank(method='ordinal', descending=True).over('col4') == 1, axis=1 ) # 条件2:行中至少有一个值大于100 cond2 = pl.any(pl.col(pivot_cols) > 100, axis=1) # 3. 应用过滤,得到最终结果 df_final_polars = p_df.filter(cond1 & cond2) print(df_final_polars)
关键说明
- Polars的
rank函数需要指定method='ordinal'来匹配Pandas默认的排名方式(避免相同值排名相同) - 处理500万行大数据时,建议改用Lazy API优化性能:将
df_p = pl.from_pandas(df)替换为df_p = pl.scan_pandas(df)(实际场景用scan_csv/scan_parquet等),最后用.collect()触发计算 - 透视后的列通过列表推导式自动获取,避免硬编码列名,提升代码通用性
内容的提问来源于stack exchange,提问作者buggsbunny4
相关产品推荐
相关产品推荐

