如何不遍历行,基于Polars DataFrame的list列生成关联分组列
高效生成Polars DataFrame的关联分组列
m_group 需求说明
现有带分组的Polars DataFrame,需新增m_group列存储组内记录的关联标识,格式为{source}_{数字}。核心规则:
- 仅
source==1的记录需要映射到组内其他所有source,因此这类记录会生成多个关联项 - 其他
source的记录仅需关联source==1,因此只有一个关联项 - 示例对应关系:
source==1的m_group=5_3,对应source==5的m_group=1_3
当前通过遍历DataFrame行生成关联列表,再执行explode()的方式效率极低,以下是优化方案。
现有低效实现
import polars as pl # 原始遍历实现(效率低) m_group_list = [] for row in df.iter_rows(named=True): if row["source"] == 1: # 查找同组内所有非1的source other_sources = df.filter( pl.col("group_id") == row["group_id"], pl.col("source") != 1 )["source"].unique() m_groups = [f"{s}_{row['seq_num']}" for s in other_sources] m_group_list.append(m_groups) else: m_group_list.append([f"1_{row['seq_num']}"]) df = df.with_columns(pl.Series(name="m_group", values=m_group_list)).explode("m_group")
示例数据
# 测试用输入DataFrame df = pl.DataFrame({ "group_id": [1, 1, 1, 2, 2], "source": [1, 5, 6, 1, 7], "seq_num": [3, 3, 3, 5, 5] })
预期输出
shape: (6, 4) ┌──────────┬────────┬────────┬─────────┐ │ group_id ┆ source ┆ seq_num┆ m_group │ │ --- ┆ --- ┆ --- ┆ --- │ │ i64 ┆ i64 ┆ i64 ┆ str │ ├──────────┼────────┼────────┼─────────┤ │ 1 ┆ 1 ┆ 3 ┆ 5_3 │ │ 1 ┆ 1 ┆ 3 ┆ 6_3 │ │ 1 ┆ 5 ┆ 3 ┆ 1_3 │ │ 1 ┆ 6 ┆ 3 ┆ 1_3 │ │ 2 ┆ 1 ┆ 5 ┆ 7_5 │ │ 2 ┆ 7 ┆ 5 ┆ 1_5 │ └──────────┴────────┴────────┴─────────┘
优化后的高效实现
完全基于Polars向量化操作,避免Python级别的行遍历,性能提升显著:
import polars as pl # 1. 按group_id聚合,获取组内非1的source列表、组内统一的seq_num group_meta = df.group_by("group_id").agg( pl.col("source").filter(pl.col("source") != 1).unique().alias("other_sources"), pl.col("seq_num").first().alias("seq_num") ) # 2. 处理source==1的记录:交叉连接到组内其他source生成多关联项 source1_result = df.filter(pl.col("source") == 1).join( group_meta, on="group_id", how="left" ).explode("other_sources").with_columns( pl.format("{}_{}", pl.col("other_sources"), pl.col("seq_num_right")).alias("m_group") ).drop("other_sources", "seq_num_right") # 3. 处理非source==1的记录:直接生成单关联项 non_source1_result = df.filter(pl.col("source") != 1).with_columns( pl.format("1_{}", pl.col("seq_num")).alias("m_group") ) # 4. 合并结果并排序 final_df = pl.concat([source1_result, non_source1_result]).sort("group_id", "source")
补充:新增m_status列(可选)
如果需要标记记录的关联类型,可添加如下代码:
final_df = final_df.with_columns( pl.when(pl.col("source") == 1).then("multiple").otherwise("single").alias("m_status") )
内容的提问来源于stack exchange,提问作者barrelquentin997
相关产品推荐
相关产品推荐

