如何高效从PySpark DataFrame的Amount列提取每条记录的最大值?
高效实现PySpark提取字符串列表中的最大值
你的原代码效率极低的核心原因是反复调用collect()把全量DataFrame拉到Driver节点处理,完全浪费了Spark的分布式计算能力——Spark设计就是要在Executor节点并行处理数据,而不是把所有数据拉到单节点循环。
下面是两种无需循环的高效实现方式,充分利用PySpark原生函数:
方法一:兼容所有Spark版本(拆分+展开+分组取最大)
如果你的Spark版本低于2.4,可以用这种方法:
from pyspark.sql import functions as F # 1. 把Amount列按逗号拆分成数组,再展开为多行记录 df_exploded = df_testing.withColumn("amount_item", F.explode(F.split(F.col("Amount"), ","))) # 2. 将拆分出的字符串转为数值类型(根据实际数据选double或int) df_exploded = df_exploded.withColumn("amount_item", F.col("amount_item").cast("double")) # 3. 按Id分组,提取每组的最大值 df_max = df_exploded.groupBy("Id").agg(F.max("amount_item").alias("MaxValue")) # 4. 和原表关联,得到包含原数据和最大值的结果 dfnew = df_testing.join(df_max, on="Id", how="inner") dfnew.show(truncate=False)
方法二:Spark 2.4+ 简洁实现(数组直接取最大)
Spark 2.4及以上支持array_max函数,可以直接对拆分后的数组取最大值,一步到位:
from pyspark.sql import functions as F dfnew = df_testing.withColumn( "MaxValue", # 先拆分字符串为数组,转数值类型,再取数组最大值 F.array_max(F.split(F.col("Amount"), ",").cast("array<double>")) ) dfnew.show(truncate=False)
关键说明
- 两种方法都是在Executor节点分布式并行处理,不会把全量数据拉到Driver,处理25000条数据只会需要几秒时间。
- 如果Amount中的元素是整数,把
cast("array<double>")改成cast("array<int>")即可。
内容的提问来源于stack exchange,提问作者S. Hasan
相关产品推荐
相关产品推荐

