Polars扩展product列行后,仅保留首行price_paid原值的实现
Polars拆分product列后仅首行保留price_paid原值的实现方案
现有Polars代码通过自定义函数将product列拆分为月度区间列表,再用explode扩展行,但price_paid列每行都保留原始值,导致总和错误。需求是:每个原始行对应的第一个扩展行保留price_paid原值,其余扩展行设为0,qty列保持不变。
原始代码:
import polars as pl import datetime as dt from dateutil.relativedelta import relativedelta def get_3_month_splits(product: str) -> list[str]: front, start_dt, total_m = product.rsplit('.', 2) start_dt = dt.datetime.strptime(start_dt, '%Y%m') total_m = int(total_m) return [f'{front}.{(start_dt+relativedelta(months=m)).strftime("%Y%m")}.3' for m in range(0, total_m, 3)] df = pl.DataFrame({ 'product': ['CHECK.GB.202403.12', 'CHECK.DE.202506.6', 'CASH.US.202509.12'], 'qty': [10, -20, 50], 'price_paid': [1400, -3300, 900], }) print(df.with_columns(pl.col('product').map_elements(get_3_month_splits, return_dtype=pl.List(str))).explode('product'))
实现思路
给每个原始行添加唯一标识,explode后按该标识分组,通过行号判断是否为组内首行,仅首行保留price_paid原值,其余置0。
完整代码
import polars as pl import datetime as dt from dateutil.relativedelta import relativedelta def get_3_month_splits(product: str) -> list[str]: front, start_dt, total_m = product.rsplit('.', 2) start_dt = dt.datetime.strptime(start_dt, '%Y%m') total_m = int(total_m) return [f'{front}.{(start_dt+relativedelta(months=m)).strftime("%Y%m")}.3' for m in range(0, total_m, 3)] df = pl.DataFrame({ 'product': ['CHECK.GB.202403.12', 'CHECK.DE.202506.6', 'CASH.US.202509.12'], 'qty': [10, -20, 50], 'price_paid': [1400, -3300, 900], }) result = ( df # 给每个原始行添加唯一ID,用于后续分组判断 .with_columns(original_row_id=pl.int_range(0, pl.count())) # 拆分product为列表 .with_columns(pl.col('product').map_elements(get_3_month_splits, return_dtype=pl.List(str))) # 扩展行 .explode('product') # 按原始行ID分组,给每组内的行编号 .with_row_number('row_in_group', by='original_row_id') # 判断是否为组内首行,首行保留原值,否则设为0 .with_columns( price_paid=pl.when(pl.col('row_in_group') == 1) .then(pl.col('price_paid')) .otherwise(0) ) # 移除临时列 .drop('original_row_id', 'row_in_group') ) print(result)
关键说明
original_row_id:生成原始行的唯一标识,确保explode后同一原始行的扩展行能被正确分组。with_row_number('row_in_group', by='original_row_id'):在每个原始行的分组内生成行号,用于判断是否为首行。pl.when().then().otherwise():条件赋值逻辑,仅组内首行保留price_paid原值,其余行置为0。
内容的提问来源于stack exchange,提问作者Phil-ZXX
相关产品推荐
相关产品推荐

