pyspark query to find the difference in column value for corresponding values in another column

Viewed 367

I am pyspark noob and trying to build a logic for the below scenario:

When the col1 has value 12 get the col2 value, and find the difference with the col2 value for the next value of col1 as value 1.

Input

For example: for the first two rows having cycle of 12 and 1, when col1 is 12 then col2 has value 11 and when col1 is 1 then col2 has value 18, difference would be 11-18 (-7 should be added in a new data frame with the corresponding timestamp value at col1==1).

The dataset is a cycle of col1 having values 12 and 1. Sometimes either col1 values 12 or 1 is repeated multiple times, in that case col2 value corresponding to the last occurrence of 12 and the first occurrence 1 value should be considered.

Expected output:

Output

Can you please help me achieve this.

1 Answers

This is some type of Gaps and Islands problem that you can solve using Window functions. First, you need to identify the "islands" (when col1 changes to 12) by using a conditional cumulative sum:

from pyspark.sql import functions as F, Window

df = spark.createDataFrame([
    ("2022-01-01T00:15:39.56", 12, 11), ("2022-01-01T00:20:37.20", 1, 18), ("2022-01-01T00:20:37.20", 12, 9),
    ("2022-01-01T00:20:37.21", 1, 6), ("2022-01-01T00:20:37.22", 12, 7), ("2022-01-01T00:21:37.21", 1, 8),
    ("2022-01-01T00:22:37.22", 1, 9), ("2022-01-01T00:22:39.22", 1, 6), ("2022-01-01T00:22:47.22", 1, 7),
    ("2022-01-01T00:23:37.25", 12, 18), ("2022-01-01T00:23:39.25", 12, 9), ("2022-01-01T00:23:39.26", 12, 7),
    ("2022-01-01T00:23:39.27", 1, 8), ], ["timestamp", "col1", "col2"]
)

w = Window.orderBy("timestamp", F.col("col1"))

df = df.withColumn(
    "group",
    F.sum(F.when(F.col("col1") == 12, 1)).over(w)
)

Now, using first function over a Window partitioned by the created column group, you can subtract value of first col2 corresponding to col1=1 from value of first col2 corresponding to col1=12, like this:

w2 = Window.partitionBy("group").orderBy("timestamp")

result = df.withColumn(
    "diff",
    F.first(F.when(F.col("col1") == 12, F.col("col2")), True).over(w2) -
    F.first(F.when(F.col("col1") == 1, F.col("col2")), True).over(w2)
).filter("col1 = 1").groupBy("group").agg(
    F.min(F.col("timestamp")).alias("timestamp"),
    F.first("diff", True).alias("diff")
).drop("group")

result.show()

#+----------------------+----+
#|timestamp             |diff|
#+----------------------+----+
#|2022-01-01T00:20:37.20|-7  |
#|2022-01-01T00:20:37.21|3   |
#|2022-01-01T00:21:37.21|-1  |
#|2022-01-01T00:23:39.27|-1  |
#+----------------------+----+
Related