Polars分组聚合优化:如何避免重复使用同一过滤器
优化方案:避免GroupBy中重复使用过滤器
在Polars里,你当前的写法会对同一列重复应用相同过滤器,确实会增加不必要的计算开销。这里有两种高效的优化方式:
方式一:对过滤后的列一次性计算多个聚合量
利用Polars的链式调用,对过滤后的列一次性生成所有需要的分位数统计量,这样过滤器只会执行一次,而非每次计算分位数都重新过滤:
filter_required = ((pl.col('sugar_level') > 40) & (pl.col('weight') >= 30)) other_filter = ((pl.col('weight') < 60)) groups = ( df.group_by('occupation') .agg( # 一次性处理filter_required对应的BMI分位数 pl.col('bmi').filter(filter_required).agg([ pl.quantile(0.10).alias('q10'), pl.quantile(0.25).alias('q25'), pl.quantile(0.50).alias('q50'), pl.quantile(0.75).alias('q75'), pl.quantile(0.90).alias('q90') ]), # 一次性处理other_filter对应的血糖分位数 pl.col('sugar_level').filter(other_filter).agg([ pl.quantile(0.25).alias('q25_sugar'), pl.quantile(0.50).alias('q50_sugar'), pl.quantile(0.75).alias('q75_sugar') ]) ) .unnest(['bmi', 'sugar_level']) # 展开嵌套的结构体列 )
这种写法让每个过滤器只执行一次,随后在过滤后的数据集上计算所有需要的分位数,大幅减少重复计算的开销。
方式二:分组内提前过滤后统一计算(进阶)
如果过滤器逻辑更复杂,也可以通过group_by+apply的方式,在每个分组内只执行两次过滤操作,再统一计算统计量:
filter_required = ((pl.col('sugar_level') > 40) & (pl.col('weight') >= 30)) other_filter = ((pl.col('weight') < 60)) def process_group(group_df): # 每个分组内仅过滤一次BMI数据 filtered_bmi = group_df.filter(filter_required)['bmi'] # 每个分组内仅过滤一次血糖数据 filtered_sugar = group_df.filter(other_filter)['sugar_level'] return pl.DataFrame({ 'q10': filtered_bmi.quantile(0.10), 'q25': filtered_bmi.quantile(0.25), 'q50': filtered_bmi.quantile(0.50), 'q75': filtered_bmi.quantile(0.75), 'q90': filtered_bmi.quantile(0.90), 'q25_sugar': filtered_sugar.quantile(0.25), 'q50_sugar': filtered_sugar.quantile(0.50), 'q75_sugar': filtered_sugar.quantile(0.75) }) groups = df.group_by('occupation').apply(process_group)
这种方式适合需要在分组内做更多自定义处理的场景,同样避免了重复过滤的问题。
额外优化:测试数据生成
你原有的测试数据生成用Python循环效率较低,可以改用Polars内置方式直接生成,速度提升明显:
import polars as pl import random df = pl.DataFrame({ 'occupation': pl.Series([random.choice(['Engineer', 'Doctor', 'Teacher', 'Artist']) for _ in range(1000000)]), 'weight': pl.Series(random.uniform(50, 100) for _ in range(1000000)), 'height': pl.Series(random.uniform(4.5, 7) for _ in range(1000000)) }).with_columns( bmi=pl.col('weight') / (pl.col('height') ** 2), sugar_level=pl.Series(random.uniform(70, 150) for _ in range(1000000)) )
内容的提问来源于stack exchange,提问作者r ram
相关产品推荐
相关产品推荐

