基于ID的PySpark多列关联查询及cat值统计需求
PySpark DataFrame 分类统计需求实现
原始数据
+-------------+----------+------+ | id | num |cat | +-------------+----------+------+ | 00111| 50012| a| | 00111| 10131| a| | 00111| 11001| b| | 10131| 71010| a| | 10131| 60010| c| | 11001| 53420| z| | 11001| 20011| a| | 11001| 00000| q| | 13403| 33001| a| | 13403| 10023| a| | 50012| 00111| a| +-------------+----------+------+
需求说明
- 针对每个唯一
id,提取其对应的所有num值 - 查找
id与上述num值匹配的所有行 - 统计这些行中
cat="a"的出现次数 - 将统计结果关联回原DataFrame的每一行,对应到所属的
id组
示例过程(以id=00111为例)
- 提取
id=00111对应的num列表:[50012, 10131, 11001] - 筛选
id属于该列表的行:
+-------------+----------+------+ | id | num |cat | +-------------+----------+------+ | 10131| 71010| a| | 10131| 60010| c| | 11001| 53420| z| | 11001| 20011| a| | 11001| 00000| q| | 50012| 00111| a| +-------------+----------+------+
- 统计
cat="a"的次数:3次
最终输出结果
+-------------+----------+------+------+ | id | num |cat |count_a| +-------------+----------+------+------+ | 00111| 50012| a| 3| | 00111| 10131| a| 3| | 00111| 11001| b| 3| | 10131| 71010| a| 1| | 10131| 60010| c| 1| | 11001| 53420| z| 0| | 11001| 20011| a| 0| | 11001| 00000| q| 0| | 13403| 33001| a| 0| | 13403| 10023| a| 0| | 50012| 00111| a| 2| +-------------+----------+------+------+
实现代码
from pyspark.sql import SparkSession from pyspark.sql import functions as F # 初始化SparkSession spark = SparkSession.builder.appName("cat_count").getOrCreate() # 创建原始DataFrame data = [ ("00111", "50012", "a"), ("00111", "10131", "a"), ("00111", "11001", "b"), ("10131", "71010", "a"), ("10131", "60010", "c"), ("11001", "53420", "z"), ("11001", "20011", "a"), ("11001", "00000", "q"), ("13403", "33001", "a"), ("13403", "10023", "a"), ("50012", "00111", "a") ] df = spark.createDataFrame(data, ["id", "num", "cat"]) # 按id分组,收集对应的num列表 id_num_groups = df.groupBy("id").agg(F.collect_list("num").alias("num_list")) # 关联原DataFrame,保留每个id对应的num列表 df_with_num_list = df.join(id_num_groups, on="id", how="left") # 将原始DataFrame重命名,避免关联冲突 df_alias = df.withColumnRenamed("id", "matched_id").withColumnRenamed("num", "matched_num").withColumnRenamed("cat", "matched_cat") # 关联找出所有id在num_list中的行 df_matched = df_with_num_list.join(df_alias, F.array_contains(df_with_num_list.num_list, df_alias.matched_id), how="left") # 按原id分组,统计cat="a"的次数 count_a_df = df_matched.groupBy("id").agg(F.sum(F.when(F.col("matched_cat") == "a", 1).otherwise(0)).alias("count_a")) # 关联回原DataFrame,得到最终结果 final_df = df.join(count_a_df, on="id", how="left") # 展示结果 final_df.show()
代码说明
- 分组收集num列表:通过
groupBy和collect_list,为每个id收集对应的所有num值形成列表。 - 关联原表:将分组结果关联回原DataFrame,让每一行都携带所属
id的num列表。 - 匹配关联:使用
array_contains判断原始DataFrame中的id是否在当前行的num列表中,筛选符合条件的行。 - 统计次数:按原
id分组,用sum和when统计matched_cat="a"的次数。 - 最终关联:将统计结果关联回原DataFrame,得到每一行对应的
count_a值。
内容的提问来源于stack exchange,提问作者user18373817
相关产品推荐
相关产品推荐

