如何用PySpark DataFrame实现数组列的差分运算?
用PySpark原生函数实现数组差分计算
你完全不需要依赖UDF来实现这个需求,PySpark的原生高阶函数就能轻松搞定,而且性能比UDF好很多——毕竟UDF需要额外的序列化/反序列化操作,会拖慢处理速度。
你的需求是对数组的相邻元素做差(即计算arr[n+1] - arr[n]),我们可以用zip_with结合slice函数来实现,完全贴合Spark的原生编程范式:
完整实现代码
from pyspark.sql import functions as F # 假设你的原始DataFrame名为df df = df.withColumn( "col_new", F.zip_with( # 截取原数组从第2个元素到末尾的部分 F.slice(F.col("col"), 2, F.size(F.col("col"))), # 截取原数组从第1个元素到倒数第2个的部分 F.slice(F.col("col"), 1, F.size(F.col("col")) - 1), # 对每一对元素执行相减操作 lambda x, y: x - y ) ) df.show(20, False)
代码细节解释
slice(col, start, length):Spark的数组截取函数,注意数组索引从1开始计数:F.slice(F.col("col"), 2, F.size(F.col("col"))):把原数组从第2个元素截取到末尾,比如原数组[1,4,3]会变成[4,3]F.slice(F.col("col"), 1, F.size(F.col("col")) - 1):把原数组从第1个元素截取到倒数第2个,比如原数组[1,4,3]会变成[1,4]
zip_with(arr1, arr2, func):将两个数组按位置配对,对每一对元素应用指定函数。这里我们让第一个数组的元素减去第二个数组的对应元素,正好得到arr[n+1]-arr[n]的结果。
运行后你会得到期望的输出:
+----------+----------+ |col |col_new | +----------+----------+ |[1, 4, 3] |[3, -1] | |[1, 5, 11]|[4, 6] | |[1, 3, 3] |[2, 0] | |[1, 4, 3] |[3, -1] | |[1, 6, 3] |[5, -3] | |[1, 1, 3] |[0, 2] | +----------+----------+
另一种实现方式(PySpark 3.1+可用)
如果你使用的是PySpark 3.1及以上版本,还可以用带索引的transform函数实现,写法更直观:
df = df.withColumn( "col_new", F.transform( # 生成从1到数组长度-1的索引序列 F.sequence(F.lit(1), F.size(F.col("col")) - 1), # 遍历每个索引,计算当前元素与前一个元素的差值 lambda i: F.col("col")[i] - F.col("col")[i-1] ) )
两种方法都完全基于Spark原生能力,避免了UDF的性能损耗,也更易于维护。
内容的提问来源于stack exchange,提问作者Boendal
相关产品推荐
相关产品推荐

