在PySpark中处理ArrayType类型的列时,一个非常常见的需求是:不仅要找出数组中的最大值,还要知道这个最大值对应的索引位置,甚至根据这个索引去另一个同长度的数组里取出对应的元素。比如一个列表存了每天的销量,另一个列表存了对应的日期,我们想直接得到销量最高的那天是哪一天。Spark本身没有一个现成的函数叫argmax,但通过组合几个内置函数,完全可以高效地实现这个效果,下面详细介绍几种方案。

方案一:array_max结合array_position获取最大值和索引
最直接的思路是分两步走:先用array_max拿到数组中的最大值,再用array_position查出这个值第一次出现的位置索引。这两个函数都从PySpark 3.1版本开始可以直接在DataFrame API中使用,如果版本较低,也可以通过F.expr调用对应的SQL函数,效果完全一样。
需要注意的是,array_position返回的是从1开始的位置,而Python的索引习惯是从0开始,所以如果后续要用这个索引去取别的数组元素,记得减1。下面的例子演示了完整的用法:
from pyspark.sql import SparkSession
from pyspark.sql import functions as F
spark = SparkSession.builder.appName("array_max_demo").getOrCreate()
df = spark.createDataFrame(
[(1, [10, 25, 7, 32, 18]), (2, [5, 9, 40, 12])],
["id", "values"]
)
result = df.withColumn("max_value", F.array_max("values")) \
.withColumn("max_index", F.array_position("values", F.array_max("values")) - 1)
result.show(truncate=False)
这个方案的优点是代码简洁、可读性好,全程使用内置函数,性能优于UDF。缺点是当数组中存在多个相同的最大值时,只能拿到第一个出现的位置,如果业务上需要所有最大值的位置,就要换用其他方法。
方案二:用transform和filter处理更复杂的场景
如果需要找出所有等于最大值的索引,可以借助高阶函数transform先生成一个带索引的辅助数组,再用filter过滤出值为最大值的那些元素。这种写法在Spark 2.4及以上版本可用,灵活性很高,适合处理更复杂的需求。
具体做法是:利用transform遍历数组时可以访问的
from pyspark.sql import functions as F
df2 = df.withColumn(
"max_value", F.array_max("values")
).withColumn(
"all_max_indexes",
F.expr("""
filter(
transform(values, (x, i) -> struct(i as idx, x as val)),
s -> s.val = array_max(values)
)
""")
)
df2.select("id", "max_value", "all_max_indexes").show(truncate=False)
这种方案虽然写法比方案一复杂一些,但表达能力更强。比如你还可以在transform里直接把另一个数组同位置的元素打包成struct,一次筛选就同时拿到索引、最大值以及关联字段,避免了多次遍历。不过要注意高阶函数嵌套过深会影响SQL的可读性,建议在代码里加好注释,或者把逻辑封装成可复用的表达式字符串。
方案三:根据索引取出关联数组中对应位置的元素
实际业务中,找到最大值的索引往往只是中间步骤,真正的目标是根据索引去另一个数组里取对应的元素。比如values数组存销量,dates数组存日期,我们要找销量最高的日期。这时可以用element_at函数配合前面得到的索引来完成取值,注意element_at同样是1-based的索引,与array_position正好配套,不需要减1。
from pyspark.sql import functions as F
df3 = spark.createDataFrame(
[(1, [10, 25, 7, 32, 18], ["a", "b", "c", "d", "e"])],
["id", "values", "labels"]
)
df3 = df3.withColumn(
"max_value", F.array_max("values")
).withColumn(
"pos", F.array_position("values", F.array_max("values"))
).withColumn(
"label_of_max", F.element_at("labels", F.col("pos"))
)
df3.show(truncate=False)
除了这种逐步计算的方式,也可以用arrays_zip把两个数组按位置压缩成一个struct数组,再用array_max配合比较逻辑找出目标元素。两种方式各有适用场景:逐步计算的方式思路清晰、易于调试;arrays_zip的方式则更适合多个数组需要同时对齐比较的情况。无论哪种方式,都建议先确认两个数组长度一致,否则压缩或取值时可能出现null或数据错位的问题。
UDF方案及性能对比
有些开发者习惯直接写一个Python UDF,在函数内部把数组转成列表,用max和list.index找到最大值和索引。这种写法确实直观,但存在明显的性能隐患:UDF会导致数据在JVM和Python进程之间来回序列化传输,破坏了Spark的Catalyst优化,无法享受Tungsten的全流程优化,在数据量大时性能差距会非常明显。
from pyspark.sql.functions import udf
from pyspark.sql.types import StructType, StructField, IntegerType
@udf(returnType=StructType([
StructField("max_val", IntegerType()),
StructField("max_idx", IntegerType())
]))
def find_max(arr):
if not arr:
return None
m = max(arr)
return (m, arr.index(m))
df.select("id", find_max("values").alias("max_info")).show(truncate=False)
一般来说,只有在内置函数实在无法表达逻辑时才考虑UDF。如果确实需要自定义逻辑,优先选择pandas_udf(即矢量化UDF),它以Arrow批次的方式传输数据,性能比普通Python UDF好很多。另外,如果整张表的数据量不大,也可以考虑先collect到Driver端再处理,但这在生产环境中要谨慎使用,避免Driver端内存溢出。
常见坑点与注意事项
使用这些函数时有几个容易踩坑的地方值得留意。首先是null处理:如果数组本身为null,array_max会返回null,array_position在目标值为null时也会返回null,最好在计算前用coalesce或过滤把空数组处理好。其次是空数组的情况,array_max对空数组同样返回null,需要根据业务决定是过滤掉还是给默认值。
其次是版本兼容问题:array_max和array_position在DataFrame API层面需要PySpark 3.1以上,低版本可以用F.expr("array_max(values)")的方式调用,因为SQL函数本身在Spark 2.4就已支持。高阶函数transform和filter则需要Spark 2.4以上。写代码前确认集群版本,能省去不少排查时间。
最后是重复最大值的问题。如果你的数据中最大值可能出现多次,一定要明确业务语义:是取第一个、取最后一个,还是全部都要。方案一默认取第一个,如果需要全部位置,就用方案二的高阶函数写法。把这个细节在代码注释中写清楚,能避免后续维护时的理解偏差。
总结
从数组列获取最大值及其索引,在PySpark中首选内置函数组合方案:array_max加array_position简单直接,transform加filter灵活强大,再配合element_at或arrays_zip就能实现跨数组取值。UDF方案虽然可行,但性能开销大,应当作为兜底选项。掌握这些函数的组合用法,不仅能解决当前问题,也能举一反三地处理数组列的各种聚合、筛选需求。
PySpark数组列最大值索引array_max修改时间:2026-09-15 12:06:36