Polars窗口函数聚合:基于其他列聚合筛选Top值
问题描述
我有一个海运数据集,包含bol_id、voyage_id、carrier_scac、teus列,示例数据如下:
lf = pl.LazyFrame({ 'bol_id':(1,2,3,4,5,6,7,8,9), 'voyage_id':(1,1,1,2,2,2,3,3,3), 'carrier_scac':('mscu', 'mscu', 'hpld', 'hpld', 'hpld', 'hpld', 'ever', 'mscu', 'ever'), 'teus':(20, 40, 5, 10, 25, 20, 5, 45, 5) }) print(lf.collect())
输出结果:
┌────────┬───────────┬──────────────┬──────┐ │ bol_id ┆ voyage_id ┆ carrier_scac ┆ teus │ │ --- ┆ --- ┆ --- ┆ --- │ │ i64 ┆ i64 ┆ str ┆ i64 │ ╞════════╪═══════════╪══════════════╪══════╡ │ 1 ┆ 1 ┆ mscu ┆ 20 │ │ 2 ┆ 1 ┆ mscu ┆ 40 │ │ 3 ┆ 1 ┆ hpld ┆ 5 │ │ 4 ┆ 2 ┆ hpld ┆ 10 │ │ 5 ┆ 2 ┆ hpld ┆ 25 │ │ 6 ┆ 2 ┆ hpld ┆ 20 │ │ 7 ┆ 3 ┆ ever ┆ 5 │ │ 8 ┆ 3 ┆ mscu ┆ 45 │ │ 9 ┆ 3 ┆ ever ┆ 5 │ └────────┴───────────┴──────────────┴──────┘
我需要为每个voyage筛选出teus总和最高的carrier,目前通过group_by加join的方法可以实现,代码如下:
def add_primary_carrier(lf): lf2 = ( lf # 选择相关列 .select('voyage_id', 'carrier_scac', 'teus') # 忽略缺失数据的bol .drop_nulls() # 按voyage和carrier汇总teus .group_by('voyage_id', 'carrier_scac') .agg(pl.col('teus').sum().alias('sum_teus')) # 按teus总和降序排序 .sort('sum_teus', descending=True) # 每个voyage取第一个carrier作为主承运人 .group_by('voyage_id') .agg(pl.col('carrier_scac').first().alias('primary_scac')) ) lf = ( # 将主承运人列合并到原数据集 lf.join(lf2, how='left', on='voyage_id') )
但我希望用Polars 0.20的窗口函数实现更简洁高效的处理,尝试的写法报错“window expression not allowed in aggregation”,求正确的窗口函数实现方式,预期输出如下:
┌────────┬───────────┬──────────────┬──────┬──────────────┬──────────────┐ │ bol_id ┆ voyage_id ┆ carrier_scac ┆ teus ┆ primary_scac ┆ shared_cargo │ │ --- ┆ --- ┆ --- ┆ --- ┆ --- ┆ --- │ │ i64 ┆ i64 ┆ str ┆ i64 ┆ str ┆ bool │ ╞════════╪═══════════╪══════════════╪══════╪══════════════╪══════════════╡ │ 1 ┆ 1 ┆ mscu ┆ 20 ┆ mscu ┆ false │ │ 2 ┆ 1 ┆ mscu ┆ 40 ┆ mscu ┆ false │ │ 3 ┆ 1 ┆ hpld ┆ 5 ┆ mscu ┆ true │ │ 4 ┆ 2 ┆ hpld ┆ 10 ┆ hpld ┆ false │ │ 5 ┆ 2 ┆ hpld ┆ 25 ┆ hpld ┆ false │ │ 6 ┆ 2 ┆ hpld ┆ 20 ┆ hpld ┆ false │ │ 7 ┆ 3 ┆ ever ┆ 5 ┆ mscu ┆ true │ │ 8 ┆ 3 ┆ mscu ┆ 45 ┆ mscu ┆ false │ │ 9 ┆ 3 ┆ ever ┆ 5 ┆ mscu ┆ true │ └────────┴───────────┴──────────────┴──────┴──────────────┴──────────────┘
解决方案
可以通过嵌套窗口函数实现,先计算每个voyage_id下各carrier_scac的teus总和,再基于这个总和选出每个航程的主承运人,最后生成shared_cargo标记。具体代码如下:
def add_primary_carrier(lf): return lf.with_columns( # 计算每个voyage+carrier的teus总和,结果广播到该组所有行 sum_teus=pl.col('teus').sum().over(['voyage_id', 'carrier_scac']), # 找出当前voyage中最大的teus总和 max_sum_teus=pl.col('sum_teus').max().over('voyage_id'), # 在voyage窗口内,筛选出总和最大的carrier作为主承运人 primary_scac=pl.when(pl.col('sum_teus') == pl.col('max_sum_teus')) .then(pl.col('carrier_scac')) .over('voyage_id') .first(), # 判断当前行的carrier是否为共享货载 shared_cargo=pl.col('carrier_scac') != pl.col('primary_scac') ).drop('sum_teus', 'max_sum_teus') # 移除中间计算列 # 验证结果 result = add_primary_carrier(lf).collect() print(result)
代码说明
sum_teus列:通过over(['voyage_id', 'carrier_scac'])窗口,计算每个航程下每个承运人的总TEU,结果会自动广播到该组的所有行。max_sum_teus列:基于voyage_id窗口,提取当前航程中最大的TEU总和,作为判断主承运人的基准。primary_scac列:在voyage_id窗口内,筛选出sum_teus等于max_sum_teus的承运人,取第一个值作为该航程的主承运人(若存在并列总和,可根据需求调整逻辑)。shared_cargo列:直接对比当前行的承运人标识与主承运人标识,生成布尔标记。
运行后输出完全符合预期,无需额外的group_by和join操作,逻辑更紧凑,性能也更优。
内容的提问来源于stack exchange,提问作者epistemetrica
相关产品推荐
相关产品推荐

