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

PySpark按单列分组获取其他列原始值列表的实现方法

PySpark分组收集原始列值列表解决方案

当然可以!这两个需求本质上都是分组后获取指定列的原始值集合,PySpark提供了非常直接的实现方式,核心就是用collect_list(保留重复值)或collect_set(去重,获取唯一可能值)函数,下面我给你一步步演示:

1. 分组后获取单列原始值列表

比如你要按某列分组,收集另一列的所有原始值,直接用groupBy配合agg和collect_list即可。举个实例:

首先创建一个示例DataFrame:

from pyspark.sql import SparkSession
from pyspark.sql.functions import collect_list, collect_set

# 初始化SparkSession
spark = SparkSession.builder.appName("GroupByCollectDemo").getOrCreate()

# 构造测试数据
data = [
    ("A", 1, "X"),
    ("A", 2, "Y"),
    ("B", 3, "X"),
    ("B", 4, "Z"),
    ("B", 5, "Y")
]
df = spark.createDataFrame(data, ["Column1", "Column2", "Column3"])
df.show()

执行后输出:

+-------+-------+-------+
|Column1|Column2|Column3|
+-------+-------+-------+
|      A|      1|      X|
|      A|      2|      Y|
|      B|      3|      X|
|      B|      4|      Z|
|      B|      5|      Y|
+-------+-------+-------+

现在按Column1分组,收集Column2的所有原始值:

# 分组并收集Column2的原始值列表
result_single = df.groupBy("Column1").agg(collect_list("Column2").alias("Column2_raw_values"))
result_single.show(truncate=False)

输出结果:

+-------+----------------+
|Column1|Column2_raw_values|
+-------+----------------+
|A      |[1, 2]          |
|B      |[3, 4, 5]       |
+-------+----------------+

2. 分组后获取多列的所有可能值列表

针对你的第二个需求——按指定列分组,获取对应列的所有可能值列表,分两种情况处理:

  • 如果需要保留所有原始重复值,继续用collect_list,可以同时收集多个列:
# 按Column1分组,同时收集Column2的所有原始值
result_multi = df.groupBy("Column1").agg(collect_list("Column2").alias("Column2_all_values"))
result_multi.show(truncate=False)
  • 如果需要的是去重后的唯一可能值,改用collect_set函数:
# 按Column1分组,获取Column2的唯一可能值列表
result_distinct = df.groupBy("Column1").agg(collect_set("Column2").alias("Column2_unique_values"))
result_distinct.show(truncate=False)

输出结果:

+-------+-------------------+
|Column1|Column2_unique_values|
+-------+-------------------+
|A      |[1, 2]             |
|B      |[3, 4, 5]          |
+-------+-------------------+

如果要同时收集多列的列表,直接在agg里添加多个collect_list/collect_set即可:

# 同时收集Column2和Column3的原始值列表
result_both = df.groupBy("Column1").agg(
    collect_list("Column2").alias("Column2_values"),
    collect_list("Column3").alias("Column3_values")
)
result_both.show(truncate=False)

几点注意事项

  • 顺序问题:collect_list会尽量保留数据的原始顺序,但Spark是分布式计算框架,若数据经过shuffle,全局顺序无法严格保证。如果需要有序的列表,可以先对DataFrame排序再分组,或者用sort_array函数对收集后的列表排序:
    from pyspark.sql.functions import sort_array
    result_sorted = df.groupBy("Column1").agg(sort_array(collect_list("Column2")).alias("Column2_sorted_values"))
    
  • 内存限制:如果某个分组的数据量极大,收集后的列表可能会占用大量内存,甚至导致OOM。这种情况下建议评估业务需求,是否真的需要所有原始值,或者考虑分批处理。

内容的提问来源于stack exchange,提问作者Amogh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 08:20:54