PySpark按条件聚合数组值:新手求助指导
Hey there! Great job getting the full sum working with AGGREGATE—let's build on that to add the conditional sum for 'buy' transactions.
The Core Idea
Since your amount and transaction_code arrays are aligned (same length, elements correspond by index), we first need to pair each amount value with its matching transaction code. Then we can use AGGREGATE to only add values where the transaction code is 'buy'.
Step-by-Step Code Implementation
First, make sure you've imported Spark's functions:
from pyspark.sql import functions as F
Then, modify your existing code to add the conditional sum column (let's call it Buy_Total):
# Add the conditional sum column df2extract = df2extract.withColumn( "Buy_Total", F.expr(""" AGGREGATE( arrays_zip(amount, transaction_code), # Pair each amount with its transaction code cast(0 as float), # Initialize accumulator to 0.0 (acc, x) -> IF(x.transaction_code = 'buy', acc + x.amount, acc) # Only add if transaction is 'buy' ) """) ) # Show the result including all columns df2extract.select("amount", "transaction_code", "Total", "Buy_Total").show()
What This Does
arrays_zip(amount, transaction_code): Creates a new array where each element is a struct combining the amount and transaction code at the same index. For your sample data, this looks like:[{amount:10, transaction_code:'buy'}, {amount:20, transaction_code:'buy'}, {amount:30, transaction_code:'sell'}, ...]AGGREGATE: Iterates over this zipped array:- Starts with an accumulator (
acc) initialized to0.0. - For each paired element (
x), checks ifx.transaction_codeis 'buy'. If yes, addsx.amounttoacc; if not, leavesaccunchanged.
- Starts with an accumulator (
- The final value of
accbecomes theBuy_Totalcolumn.
Sample Output
For your example data, the result will look like this:
+--------------------+--------------------+-----+---------+ | amount| transaction_code|Total|Buy_Total| +--------------------+--------------------+-----+---------+ |[10, 20, 30, 40, ...|[buy, buy, sell, ...|210.0| 30.0| +--------------------+--------------------+-----+---------+
Notes for Scala Users
If you're working with Scala instead of Python, the logic is identical—just adjust the syntax slightly:
import org.apache.spark.sql.functions._ df2extract = df2extract.withColumn( "Buy_Total", expr(""" AGGREGATE( arrays_zip(amount, transaction_code), cast(0 as float), (acc, x) -> IF(x.transaction_code = 'buy', acc + x.amount, acc) ) """) ).select("amount", "transaction_code", "Total", "Buy_Total").show()
内容的提问来源于stack exchange,提问作者keerthi007

