PySpark SQL: How to use multiple conditions in "Window -> PartitionBy -> Range Between"

Viewed 252

Here is our test input data (85 rows):

+-------------+-------+---------+
|    auth_dttm|acc_num|tranAmnts|
+-------------+-------+---------+
|11/8/20 11:20|    123|      100|
|11/8/20 11:19|    123|      100|
|11/8/20 11:18|    123|      100|
|11/8/20 11:17|    123|      100|
|11/8/20 11:16|    123|      100|
|11/8/20 11:15|    123|      100|
|11/8/20 11:14|    123|      100|
|11/8/20 11:13|    123|      100|
|11/8/20 11:12|    123|      100|
|11/8/20 11:11|    123|      100|
|11/8/20 11:10|    123|      100|
|11/8/20 11:09|    123|      100|
|11/8/20 11:08|    123|      100|
|11/8/20 11:07|    123|      100|
|11/8/20 11:06|    123|      100|
|11/8/20 11:05|    123|      100|
|11/8/20 11:04|    123|      100|
|11/8/20 11:03|    123|      100|
|11/8/20 11:02|    123|      100|
|11/8/20 11:01|    123|      100|
|11/8/20 11:00|    123|      100|
|11/8/20 10:59|    123|      100|
|11/8/20 10:58|    123|      100|
|11/8/20 10:57|    123|      100|
|11/8/20 10:56|    123|      100|
|11/8/20 10:55|    123|      100|
|11/8/20 10:54|    123|      100|
|11/8/20 10:53|    123|      100|
|11/8/20 10:52|    123|      100|
|11/8/20 10:51|    123|      100|
|11/8/20 10:50|    123|      100|
|11/8/20 11:20|    321|    10000|
|11/8/20 11:19|    321|    10000|
|11/8/20 11:18|    321|    10000|
|11/8/20 11:17|    321|    10000|
|11/8/20 11:16|    321|    10000|
|11/8/20 11:15|    321|    10000|
|11/8/20 11:14|    321|    10000|
|11/8/20 11:13|    321|    10000|
|11/8/20 11:12|    321|    10000|
|11/8/20 11:11|    321|    10000|
|11/8/20 11:10|    321|    10000|
|11/8/20 11:09|    321|    10000|
|11/8/20 11:08|    321|    10000|
|11/8/20 11:07|    321|    10000|
|11/8/20 11:06|    321|    10000|
|11/8/20 11:05|    321|    10000|
|11/8/20 11:04|    321|    10000|
|11/8/20 11:03|    321|    10000|
|11/8/20 11:02|    321|    10000|
|11/8/20 11:01|    321|    10000|
|11/8/20 11:00|    321|    10000|
|11/8/20 10:59|    321|    10000|
|11/8/20 10:58|    321|    10000|
|11/8/20 10:57|    321|    10000|
|11/8/20 10:56|    321|    10000|
|11/8/20 10:55|    321|    10000|
|11/8/20 10:54|    321|    10000|
|11/8/20 10:53|    321|    10000|
|11/8/20 10:52|    321|    10000|
|11/8/20 10:51|    321|    10000|
|11/8/20 10:50|    321|    10000|
|11/8/20 10:49|    321|    10000|
|11/8/20 10:48|    321|    10000|
|11/8/20 10:47|    321|    10000|
|11/8/20 10:46|    321|    10000|
|11/8/20 10:45|    321|    10000|
|11/8/20 10:44|    321|    10000|
|11/8/20 10:43|    321|    10000|
|11/8/20 10:42|    321|    10000|
|11/8/20 10:41|    321|    10000|
|11/8/20 10:40|    321|    10000|
|11/8/20 10:39|    321|    10000|
|11/8/20 10:38|    321|    10000|
|11/8/20 10:37|    321|    10000|
|11/8/20 10:36|    321|    10000|
|11/8/20 10:35|    321|    10000|
|11/8/20 10:34|    321|    10000|
|11/8/20 10:33|    321|    10000|
| 9/1/20 11:18|    321|    10000|
| 7/1/20 11:18|    321|    10000|
| 5/1/20 11:18|    321|    10000|
| 3/1/20 11:18|    321|    10000|
| 1/1/20 11:18|    321|    10000|
+-------------+-------+---------+

What I am trying to do is:

  • Count and sum top 45 transactions in the last 24 hours for every transaction

I am able to do with the following approach:

  1. Self Join

  2. Add a new Time Diff column "DiffinSeconds"

  3. Filter by "DiffinSeconds > 0 and DiffinSeconds < 86400"

  4. Adding new column (row_number) using partitionBy AccNum,DateTime OrderBy DiffinSeconds

  5. Sum and count top 45 rows

    # PySpark Imports
    from pyspark.sql import SparkSession
    from pyspark.sql import functions as F
    from pyspark.sql.functions import desc,asc
    from pyspark.sql import Window
    
    # Create a Spark Session
    spark = SparkSession \
        .builder \
        .appName("test") \
        .config('spark.sql.legacy.timeParserPolicy', 'LEGACY') \
        .getOrCreate()
    
    # Read input file
    sparkDf = spark.read.csv('input.csv',header=True)
    
    # DateTime conversion (converting String date to PySpark dates)
    sparkDf = (
        sparkDf
        .withColumn('datetime', F.to_timestamp(sparkDf.auth_dttm, 'M/d/yyyy HH:mm'))
        .withColumn("date", F.to_date(F.to_timestamp(sparkDf.auth_dttm, 'M/d/yyyy HH:mm')))
        .drop('auth_dttm')
    )
    
    # Self Join on account, filter datediff which are within last 24 hours
    joinedDf = (
        sparkDf.alias('df1').join(sparkDf.alias("df2"), F.col("df2.acc_num") == F.col("df2.acc_num"), "inner")
        .withColumn('DiffInSeconds',F.unix_timestamp(F.col("df1.datetime")) - F.unix_timestamp(F.col('df2.datetime')))
        .select(
            F.col('df1.acc_num'),
            F.col('df1.datetime')
            F.col('df1.tranAmnts'),
            F.col('df2.datetime').alias('trailing_datetime'), 
            F.col('df2.tranAmnts').alias('trailing_tranAmnts'), 
            'DiffInSeconds'
        )
        .filter('DiffInSeconds >= 0 and DiffInSeconds <= 86400')
    )
    
    
    # Define a windows function to partition on account and datetime, Order DiffInSeconds
    window = Window.partitionBy("acc_num", "datetime").orderBy(asc(F.col("DiffInSeconds")))
    
    # Add a row_count column using above window
    windowDf = (
        joinedDf
        .withColumn("row_count", F.row_number().over(window))
    )
    
    # Final Aggregation and filter by row count
    aggDf = (
        windowDf
        .filter(
            "row_count <= 45"
        ).groupBy('acc_num', 'datetime')
        .agg(
          F.sum('tranAmnts').alias('sumTranAmnts'),
          F.count('tranAmnts').alias('countTranAmnts'))
    )
    
    aggDf.show()
    

Output:

+-------+-------------------+------------+--------------+
|acc_num|           datetime|sumTranAmnts|countTranAmnts|
+-------+-------------------+------------+--------------+
|    321|0020-11-08 11:18:00|    450000.0|            45|
|    321|0020-11-08 10:44:00|    120000.0|            12|
|    123|0020-11-08 10:56:00|       700.0|             7|
|    321|0020-11-08 11:09:00|    370000.0|            37|
|    123|0020-11-08 11:14:00|      2500.0|            25|
|    321|0020-11-08 10:51:00|    190000.0|            19|
|    321|0020-11-08 10:53:00|    210000.0|            21|
|    123|0020-11-08 10:55:00|       600.0|             6|
|    321|0020-07-01 11:18:00|     10000.0|             1|
|    123|0020-11-08 11:00:00|      1100.0|            11|
|    123|0020-11-08 11:19:00|      3000.0|            30|
|    321|0020-11-08 11:20:00|    450000.0|            45|
|    321|0020-11-08 11:01:00|    290000.0|            29|
|    321|0020-11-08 10:46:00|    140000.0|            14|
|    321|0020-11-08 10:34:00|     20000.0|             2|
|    321|0020-11-08 10:36:00|     40000.0|             4|
|    123|0020-11-08 11:09:00|      2000.0|            20|
|    123|0020-11-08 11:13:00|      2400.0|            24|
|    321|0020-11-08 11:17:00|    450000.0|            45|
|    321|0020-11-08 10:50:00|    180000.0|            18|
+-------+-------------------+------------+--------------+
only showing top 20 rows

My concern: This works great on small dataset, but I am pretty sure, "self join" will blow up when number of rows are in millions.

I am trying to solve this without self join, this is what I have till now:

sparkDf.registerTempTable('input')
df = spark.sql("""
    SELECT 
        acc_num, 
        tranAmnts, 
        datetime, 
        sum(tranAmnts) OVER (
            PARTITION BY acc_num 
            ORDER BY datetime 
            RANGE BETWEEN INTERVAL 24 HOURS PRECEDING AND CURRENT ROW) AS totalAmnt,
        count(tranAmnts) OVER (
            PARTITION BY acc_num 
            ORDER BY datetime 
            RANGE BETWEEN INTERVAL 24 HOURS PRECEDING AND CURRENT ROW) AS totalCount 
    from input
""")

But I am unable to figure out how to use multiple conditions in "RANGE BETWEEN", so I can specify both conditions for past 24h and take top 45.

Edit: As I haven't received an answer on how to use multiple conditions in "RANGE BETWEEN" clause, I would like to see if someone has suggestions on how can I improve the working "self-join" to make it more performant.

Thanks in Advance, Hussain Bohra

0 Answers
Related