PySpark DataFrame查找array列中最接近int列值的元素索引
PySpark获取array列中与int列数值最接近元素的索引实现方案
你之前的实现存在两个核心问题:
- 基于pandas/numpy的自定义逻辑没有封装成符合PySpark规范的UDF,无法直接作用于分布式DataFrame
array_position仅支持定位精确匹配的元素,无法匹配最接近的非等值元素
方案1:内置高阶函数实现(推荐,Spark 2.4+支持)
纯Spark原生函数实现,没有跨进程序列化开销,性能远高于UDF方案,代码如下:
from pyspark.sql import functions as F df = df.withColumn( "closest_index", F.expr(""" array_position( transform(array_col, x -> abs(x - int_col)), array_min(transform(array_col, x -> abs(x - int_col))) ) - 1 """) )
逻辑说明:
- 用
transform遍历array列所有元素,计算每个元素和int列数值的绝对差值,生成差值数组 - 用
array_min获取差值数组中的最小差值 - 用
array_position定位最小差值在差值数组中的位置,Spark数组默认从1开始计数,减1后即为所需的0起始索引
如果存在多个元素和目标值差值相同,该方案会返回第一个符合条件的元素索引,符合通用场景预期。
方案2:Python UDF实现
如果偏好更简洁的代码逻辑,且数据量不大,可以选择UDF方案:
import numpy as np from pyspark.sql import functions as F from pyspark.sql.types import IntegerType @F.udf(returnType=IntegerType()) def get_closest_idx(arr, target): return int(np.argmin(np.abs(np.array(arr) - target))) df = df.withColumn("closest_index", get_closest_idx(F.col("array_col"), F.col("int_col")))
两种方案对示例数据执行后,输出结果和你给出的预期完全一致。
内容的提问来源于stack exchange,提问作者Bjorno
相关产品推荐
相关产品推荐

