PySpark实现分组累积减法遇异常,求正确实现方案
PySpark分组累积减法实现方案
你的问题出在PySpark的列计算基于原始列值,不支持逐行迭代更新。原代码中,当rnb>1时,lag(F.col("_UNCS"))只能取到第一行的有效值,从第三行开始lag返回null,导致后续计算结果全部为空。
正确实现思路
观察预期结果可知,_uncs的本质是:分组内首行的councs值,减去从第二行到当前行所有councs的总和。我们可以通过窗口函数的累积求和实现,无需依赖lag进行迭代计算:
- 按
coitm分区,保持你原有的排序规则(desc("COCHDJ")) - 构造辅助列:首行保留原
councs值,后续行取councs的负数 - 对辅助列做窗口内的累积求和,结果即为预期的
_uncs
修正后的代码
from pyspark.sql import functions as F from pyspark.sql.window import Window # 定义窗口(注意统一字段大小写,避免原代码中大小写混用的问题) window_spec = Window.partitionBy("coitm").orderBy(F.desc("COCHDJ")) # 合并逻辑,直接生成目标列 output_df = _cost_srt.withColumn( "_uncs", F.sum( F.when(F.row_number().over(window_spec) == 1, F.col("councs")) .otherwise(-F.col("councs")) ).over(window_spec) ) output_df.show()
结果验证
运行代码后,_uncs列将完全匹配预期结果:
coitm=1110组:1000 → 900 → 850 → 820 → 800coitm=1210组:2000 → 1990 → 1790 → 1890 → 1870
内容的提问来源于stack exchange,提问作者Arpan Ghimire
相关产品推荐
相关产品推荐

