如何在PySpark DataFrame数组中获取前两大元素及其索引
Extract Top Two Largest Values with Indices from PySpark Array Column
Hey there! Let's solve your problem where you need to pull the top two largest values from an array column in PySpark, along with their original positions in the array. Here's a step-by-step solution using built-in PySpark functions (way more efficient than UDFs for big datasets):
Sample Input DataFrame
First, let's recreate your input DataFrame to work with:
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.window import Window spark = SparkSession.builder.appName("TopTwoArrayElements").getOrCreate() # Your input data data = [ ([0.27047928569511825, 0.5312608102025099, 0.19825990410237174],), ([0.06711381377029987, 0.8775456658890036, 0.05534052034069637],), ([0.10847074295048188, 0.04602848157663474, 0.8455007754728833],) ] df = spark.createDataFrame(data, ["probability"]) df.show(truncate=False)
This outputs your original DataFrame:
| probability |
|---|
| [0.27047928569511825, 0.5312608102025099, 0.19825990410237174] |
| [0.06711381377029987, 0.8775456658890036, 0.05534052034069637] |
| [0.10847074295048188, 0.04602848157663474, 0.8455007754728833] |
Solution Code
Here's how to transform this into your desired output:
- Add a unique row ID: This helps group elements back to their original rows after exploding the array.
- Explode the array with indices: Use
posexplodeto split the array into individual elements and their positions. - Rank elements by value: Use a window function to rank elements in descending order for each original row.
- Pivot to get top two values/indices: Filter for the top 2 ranks, then pivot the results into separate columns.
# Step 1: Add unique row identifier df_with_id = df.withColumn("row_id", F.monotonically_increasing_id()) # Step 2: Explode array with indices exploded_df = df_with_id.select( "row_id", "probability", F.posexplode("probability").alias("index", "value") ) # Step 3: Rank elements by descending value per row window_spec = Window.partitionBy("row_id").orderBy(F.desc("value")) ranked_df = exploded_df.withColumn("rank", F.row_number().over(window_spec)) # Step 4: Pivot to get top two values and indices result_df = ranked_df.filter(F.col("rank").isin(1, 2)) \ .groupBy("row_id", "probability") \ .pivot("rank", [1, 2]) \ .agg( F.first("value").alias("value"), F.first("index").alias("index") ) \ .select( "probability", F.col("1_value").alias("largest_1"), F.col("1_index").alias("index_1"), F.col("2_value").alias("largest_2"), F.col("2_index").alias("index_2") ) \ .orderBy("row_id") # Show the final result result_df.show(truncate=False)
Final Output
Running this code gives you exactly the result you want:
| probability | largest_1 | index_1 | largest_2 | index_2 |
|---|---|---|---|---|
| [0.27047928569511825, 0.5312608102025099, 0.19825990410237174] | 0.5312608102025099 | 1 | 0.27047928569511825 | 0 |
| [0.06711381377029987, 0.8775456658890036, 0.05534052034069637] | 0.8775456658890036 | 1 | 0.06711381377029987 | 0 |
| [0.10847074295048188, 0.04602848157663474, 0.8455007754728833] | 0.8455007754728833 | 2 | 0.10847074295048188 | 0 |
Notes
- Handling duplicates: If your array has duplicate maximum values,
row_number()will assign distinct ranks. If you want to treat duplicates as equal (e.g., two values tied for first), replacerow_number()withrank()ordense_rank()in the window function. - Performance: This approach uses PySpark's optimized built-in functions, which scale way better than Python UDFs for large distributed datasets.
内容的提问来源于stack exchange,提问作者luckcul
相关产品推荐
相关产品推荐

