PySpark按var_col_name指定列的值聚合id列的问题
问题分析
你的代码错误在于仅按name和var_col_name分组/开窗,没有考虑var_col_name所指定列的实际值。以name为b的行为例,它们的var_col_name都是next_col,但next_col的实际值分别是biggest text和ghjljkk,属于不同分组,不应聚合id。
正确实现方式
核心思路是:动态提取var_col_name指定列的值,将其加入分组/开窗的分区条件中,确保只有当name相同、var_col_name相同,且指定列的实际值也相同时,才聚合id。
方法1:分组关联法
from pyspark.sql import SparkSession from pyspark.sql import functions as F spark = SparkSession.builder.getOrCreate() data = [ ('a', 1, 'text1', 'col1', 'texts1', 'scj', 'dsiul'), ('a', 11, 'text1', 'col1', 'texts1', 'ftjjjjjjj', 'jhkl'), ('b', 2, 'bigger text', 'next_col', 'gfsajh', 'xcj', 'biggest text'), ('b', 21, 'bigger text', 'next_col', 'fghm', 'hjjkl', 'ghjljkk'), ('c', 3, 'soon', 'column', 'szjcj', 'sooner', 'sjdsk') ] df = spark.createDataFrame(data, ['name', 'id', 'txt', 'var_col_name', 'col1', 'column', 'next_col']) # 1. 动态提取var_col_name指定列的值,新增为var_col_value列 df_with_var = df.withColumn("var_col_value", F.col(F.col("var_col_name"))) # 2. 按name、var_col_name、var_col_value分组,聚合id grouped_df = df_with_var.groupBy("name", "var_col_name", "var_col_value")\ .agg(F.collect_set("id").alias("id_all")) # 3. 关联回原表,获取最终结果 df_agg = df.join(grouped_df, on=["name", "var_col_name", "var_col_value"], how="left")\ .select(df["*"], "id_all") df_agg.show(truncate=False)
方法2:窗口函数法
如果不想额外关联,也可以直接用窗口函数实现:
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.window import Window spark = SparkSession.builder.getOrCreate() data = [ ('a', 1, 'text1', 'col1', 'texts1', 'scj', 'dsiul'), ('a', 11, 'text1', 'col1', 'texts1', 'ftjjjjjjj', 'jhkl'), ('b', 2, 'bigger text', 'next_col', 'gfsajh', 'xcj', 'biggest text'), ('b', 21, 'bigger text', 'next_col', 'fghm', 'hjjkl', 'ghjljkk'), ('c', 3, 'soon', 'column', 'szjcj', 'sooner', 'sjdsk') ] df = spark.createDataFrame(data, ['name', 'id', 'txt', 'var_col_name', 'col1', 'column', 'next_col']) # 定义窗口:分区条件包含name、var_col_name,以及动态提取的指定列值 window = Window.partitionBy( "name", "var_col_name", F.col(F.col("var_col_name")) # 动态获取var_col_name对应的列值作为分区键 ) # 计算每个分组的id集合 df_agg = df.withColumn("id_all", F.collect_set("id").over(window)) df_agg.show(truncate=False)
结果验证
两种方法都会得到正确结果:
- name=a的两行,
var_col_name是col1,且col1值都是texts1,所以id_all为[1,11] - name=b的两行,
var_col_name是next_col,但next_col值分别为biggest text和ghjljkk,所以各自的id_all分别为[2]和[21] - name=c的单行,
id_all为[3]
内容的提问来源于stack exchange,提问作者viji
相关产品推荐
相关产品推荐

