PySpark DataFrame基于多列关联更新val列值的问题
解决PySpark按tests或cnty分组更新val的问题
你已经迈出了正确的第一步,但只考虑了tests分组的情况,没覆盖到cnty分组的关联判断——这就是为什么test2的val还是N的核心原因:它所属的cnty=UK分组里确实存在val=Y的记录,但你的代码没有把这个分组的条件纳入判断逻辑。
我们需要同时检查每条记录所属的tests分组和cnty分组是否存在val=Y的情况,只要其中任意一个分组满足条件,就把该记录的val改为Y。具体实现步骤如下:
步骤1:定义分组窗口并标记分组是否存在Y值
我们需要创建两个窗口分别对应tests和cnty的分组,然后用标记值来判断每个分组内是否存在Y:
- 按
tests分区,标记该分组是否有val=Y - 按
cnty分区,标记该分组是否有val=Y
步骤2:根据分组标记更新val字段
只要tests分组或cnty分组的标记显示存在Y,就将该记录的val设为Y,否则保留原有值。
完整代码实现
from pyspark.sql import Window import pyspark.sql.functions as f # 定义两个分组窗口 window_tests = Window.partitionBy('tests') window_cnty = Window.partitionBy('cnty') # 计算分组标记并更新val df_result = df.withColumn( 'tests_has_Y', # 用1标记存在Y,0标记不存在,取max判断分组内是否有Y f.max(f.when(f.col('val') == 'Y', 1).otherwise(0)).over(window_tests) ).withColumn( 'cnty_has_Y', f.max(f.when(f.col('val') == 'Y', 1).otherwise(0)).over(window_cnty) ).withColumn( 'val', # 只要任意一个分组存在Y,就把val设为Y f.when( (f.col('tests_has_Y') == 1) | (f.col('cnty_has_Y') == 1), 'Y' ).otherwise(f.col('val')) ).drop('tests_has_Y', 'cnty_has_Y') # 删除中间辅助列 df_result.show()
预期输出结果
+-----+---+---+----+ |tests|val|asd|cnty| +-----+---+---+----+ |test1| Y| 1|null| |test1| Y| 2|null| |test1| Y| 3|null| |test2| Y| 2| UK| # 这里val成功更新为Y,因为UK分组存在Y |test3| N| 4| AUS| # AUS和test3分组都无Y,保留原val |test4| Y| 5|null| | null| Y| 1| UK| +-----+---+---+----+
关键逻辑解释
tests_has_Y和cnty_has_Y是辅助列,用来标记当前记录所属的分组是否存在val=Y- 最终的
when判断逻辑完全匹配你的需求:只要任意一个关联分组存在Y,就统一更新val为Y,否则保留原值 - 最后删除辅助列,得到符合要求的干净结果
内容的提问来源于stack exchange,提问作者User12345
相关产品推荐
相关产品推荐

