PySpark基于flag列条件计算衍生sum列报错的解决方法
PySpark按flag列逐行计算实现方案
现有问题梳理
输入的测试DataFrame结构与样例数据:
+-----------+---------+------------------+----------------------+-----------+ | DATE | ID |sal | vat | flag | +-----------+---------+------------------+----------------------+-----------+ |10-may-2022| 1 | 1000.0| 12.0| 1 | |12-may-2022| 2 | 50.0 | 6.0| 1 | +-----------+---------+------------------+----------------------+-----------+
需要实现的计算规则:
- 行数据
flag值为1时,新列sum取值为sal * 2 - 行数据
flag值为2时,新列sum取值为sal * 4
你写的代码有两个核心问题:
- 语法层面:Python对缩进敏感,
if/else块下的执行语句没有缩进,会直接报语法错误 - 逻辑层面:
srcdf.select(col("flag"))返回的是一个DataFrame类型的分布式数据集,不是单个数值,根本不能直接和1做相等判断;另外你要的是逐行按flag值计算,不是判断整个表的flag是不是全等于1,用Python原生if/else根本实现不了逐行分支。
标准实现方式
逐行分支计算直接用PySpark内置的when().otherwise()条件函数就行,全程在Spark执行层完成,不需要拉取数据到本地,性能最优:
from pyspark.sql.functions import col, when df = srcdf.withColumn( "sum", when(col("flag") == 1, col("sal") * 2) .when(col("flag") == 2, col("sal") * 4) # 非1、2的flag值默认填null,有需要可以改成其他默认值 .otherwise(None) ) display(df)
特殊场景补充
如果你实际需求是先判断整张表的flag列全为某个固定值,再给全表统一计算sum列(不是逐行判断),也需要先把结果收集为标量值再做判断,不能直接拿DataFrame对象和数值比较:
# 仅适用于全表统一逻辑判断场景,逐行计算不要这么写 first_flag = srcdf.select("flag").first()[0] if first_flag == 1: df = srcdf.withColumn("sum", col("sal") * 2) else: df = srcdf.withColumn("sum", col("sal") * 4) display(df)
提醒:只要是逐行计算的场景,优先用Spark内置函数实现,不要强行把分布式数据拉到Driver端用Python原生逻辑处理,数据量上来之后会直接出现性能瓶颈甚至OOM。
内容的提问来源于stack exchange,提问作者SanjanaSanju
相关产品推荐
相关产品推荐

