I have this data frame:
df = (
spark
.createDataFrame([
[20210101, 'A', 103, "abc"],
[20210101, 'A', 102, "def"],
[20210101, 'A', 101, "def"],
[20210102, 'A', 34, "ghu"],
[20210101, 'B', 180, "xyz"],
[20210102, 'B', 123, "kqt"]
]
).toDF("txn_date", "txn_type", "txn_amount", "other_attributes")
)
Each date has multiple transactions of each of the different types. My task is to compute the standard deviation of the amount for each record (for the same type and going back 30 days).
The most obvious approach (that I tried) is to create a window based on type and include records going back to past 30 days.
days = lambda i: i * 86400
win = Window.partitionBy("txn_type").orderBy(F.col("txn_date").cast(LongType())).rangeBetween(-days(30), 0)
df = df.withColumn("stddev_last_30days", F.stddev(F.col("txn_amount")).over(win))
Since some of the transaction types have millions of transactions per day, this runs into OOM.
I tried doing it in parts (take only few records for each date at a time) but this leads to error prone calculations since standard deviation is not additive.
I also tried 'collect_set' for all records for a transaction type and date (so all amounts come in as an array in one column), but this runs into OOM as well.
I tried processing one month at a time (I need at a minimum 2 months data since I need to go back 1 month) but even that overwhelms my executors.
What would be a scalable way to solve this problem?
Notes:
In the original data, column
txn_dateis stored as long in "yyyyMMdd" format.There are other columns in the data frame that may or may not be same for each date and type. I haven't included them in the sample code for simplicity.