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

基于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为例)

  1. 提取id=00111对应的num列表:[50012, 10131, 11001]
  2. 筛选id属于该列表的行:
+-------------+----------+------+
|        id   |      num |cat   |
+-------------+----------+------+
|        10131|     71010|     a|
|        10131|     60010|     c|
|        11001|     53420|     z|
|        11001|     20011|     a|
|        11001|     00000|     q|
|        50012|     00111|     a|
+-------------+----------+------+
  1. 统计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()

代码说明

  1. 分组收集num列表:通过groupBy和collect_list,为每个id收集对应的所有num值形成列表。
  2. 关联原表:将分组结果关联回原DataFrame,让每一行都携带所属id的num列表。
  3. 匹配关联:使用array_contains判断原始DataFrame中的id是否在当前行的num列表中,筛选符合条件的行。
  4. 统计次数:按原id分组,用sum和when统计matched_cat="a"的次数。
  5. 最终关联:将统计结果关联回原DataFrame,得到每一行对应的count_a值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 10:45:00