Spark中根据索引列提取数组列元素的高效优雅实现方案
问题:根据索引数组提取数组列对应元素
我有如下两列数据:
| col_arr | col_ind |
|---|---|
| [1, 2, 3] | [0, 2] |
| [5, 1] | [1] |
希望根据col_ind中的索引提取col_arr对应位置的元素,得到如下结果中的col_val列:
| col_arr | col_ind | col_val |
|---|---|---|
| [1, 2, 3] | [0, 2] | [1, 3] |
| [5, 1] | [1] | [1] |
我最初考虑使用UDF,但感觉有些冗余:
@udf def sub_select(arr, inds): if (arr is not None) and (inds is not None): return [arr[ind] for ind in inds]
我还考虑结合expr动态使用array_position函数,但不清楚如何适配col_ind的长度:
F.expr("array_position(col_arr, array_position(col_ind, 0))")
补充场景:
- 假设
col_ind的最大长度较小(例如最多为5)。 - 若存在多个数组列(
col_arr1、col_arr2、col_arr3),但只有一个索引列col_ind该如何处理?
最优实现方案
方案一:使用Spark内置函数transform(通用高效)
Spark 2.4+支持的transform高阶函数是最优雅高效的选择,无需自定义UDF,直接对col_ind中的每个索引做映射提取:
import pyspark.sql.functions as F df = df.withColumn( "col_val", F.transform( F.col("col_ind"), lambda idx: F.col("col_arr")[idx] ) )
该写法原生支持空值处理:若col_arr或col_ind为空,结果自动返回空数组,无需额外判断逻辑,性能远优于自定义UDF。
方案二:针对col_ind最大长度较小的场景(补充场景1)
如果col_ind长度固定或最大长度有限(比如最多5个元素),可以结合element_at和array手动拼接,小数据量下性能表现稳定:
# 假设col_ind最多包含5个元素 df = df.withColumn( "col_val", F.array( F.when(F.size(F.col("col_ind")) >= 1, F.col("col_arr")[F.element_at(F.col("col_ind"), 1)]), F.when(F.size(F.col("col_ind")) >= 2, F.col("col_arr")[F.element_at(F.col("col_ind"), 2)]), F.when(F.size(F.col("col_ind")) >= 3, F.col("col_arr")[F.element_at(F.col("col_ind"), 3)]), F.when(F.size(F.col("col_ind")) >= 4, F.col("col_arr")[F.element_at(F.col("col_ind"), 4)]), F.when(F.size(F.col("col_ind")) >= 5, F.col("col_arr")[F.element_at(F.col("col_ind"), 5)]) ).filter(F.col("col_val").isNotNull()) )
通过filter过滤空值,确保结果数组只保留有效提取的元素。
处理多数组列+单索引列的场景(补充场景2)
多个数组列共享同一索引列时,直接复用transform逻辑即可,代码简洁易维护:
# 逐个处理多数组列 df = df.withColumn("col_val1", F.transform(F.col("col_ind"), lambda idx: F.col("col_arr1")[idx]))\ .withColumn("col_val2", F.transform(F.col("col_ind"), lambda idx: F.col("col_arr2")[idx]))\ .withColumn("col_val3", F.transform(F.col("col_ind"), lambda idx: F.col("col_arr3")[idx]))
如果数组列数量较多,用循环简化代码:
arr_cols = ["col_arr1", "col_arr2", "col_arr3"] for arr_col in arr_cols: df = df.withColumn( f"{arr_col}_val", F.transform(F.col("col_ind"), lambda idx: F.col(arr_col)[idx]) )
内容的提问来源于stack exchange,提问作者Gal Novich
相关产品推荐
相关产品推荐

