如何在PySpark中按分组过滤DataFrame 筛选含至少1个flag=1的家庭行
PySpark实现按家庭过滤的两种方案
需求说明
现有包含Household、Person、flag三列的PySpark DataFrame,需要过滤掉所有没有flag=1记录的家庭的全部行,仅保留至少存在一条flag=1记录的家庭的所有数据。
实现方法
方法1:窗口函数实现(写法简洁,适合中小数据集)
通过窗口函数按家庭分组取flag的最大值,只要分组内有一条flag=1,最大值就为1,直接过滤最大值为1的分组所有行即可:
from pyspark.sql import Window from pyspark.sql.functions import max, col # 定义按家庭分组的窗口 window_household = Window.partitionBy("Household") # 新增辅助列标记当前家庭是否存在flag=1的记录 df = df.withColumn("has_valid_flag", max(col("flag")).over(window_household)) # 过滤保留符合要求的行,删除辅助列 df = df.filter(col("has_valid_flag") == 1).drop("has_valid_flag")
方法2:分组Join实现(性能更优,适合超大规模数据集)
先提取所有存在flag=1的家庭ID,再和原表做内连接,仅保留符合要求的家庭的全部数据:
from pyspark.sql.functions import col # 提取所有符合要求的家庭ID(去重避免重复关联) valid_households = df.filter(col("flag") == 1).select("Household").distinct() # 内连接过滤得到最终结果 df = df.join(valid_households, on="Household", how="inner")
注意事项
如果你的flag列是字符串类型,将过滤条件中的1改为"1"即可。
内容的提问来源于stack exchange,提问作者Olivander
相关产品推荐
相关产品推荐

