You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何不遍历行,基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.20 23:06:09