Speeding up a custom aggregate for window function

Viewed 56

I am looking for alternatives to this code which would run faster

inp = spark.createDataFrame([
    ["1", "A",  7, 2],
    ["1", "A", 14, 3],
    ["1", "A", 35, 2],
    ["1", "A", 42, 3],
    ["1", "B", 14, 1],
    ["1", "B", 84, 2],
    ["2", "A", 14, 1],
    ["2", "A", 21, 1],
    ["2", "A", 21, 2],
  ], schema=["id","grp","elap","lbl"])
inp.show()

Desired output is

thresh = 2

@udf(returnType=IntegerType())
def best_label(lst):
  ctr = Counter(lst)
  for N in range(thresh,0,-1):
    tmp = [k for k,v in ctr.items() if v>=N]
    if len(tmp)>0:
      return max(tmp)

w = W.partitionBy("id","grp").orderBy("elap").rangeBetween(-49,0)
out = (
  inp.withColumn("lbl", F.collect_list("lbl").over(w))
    .withColumn("lbl2", best_label(F.col("lbl"))).distinct()
    .orderBy("id","grp","elap")
  )
out.show()

I am running this on a dataframe with 300 million rows and it takes about 8 mins.

1 Answers

Without using UDF it should perform much better. You can use this trick by first counting the values of lbl over the same window frame, then get the max by count taking advantage of struct ordering:

from pyspark.sql import functions as F, Window

w = Window.partitionBy("id", "grp", "lbl").orderBy("elap").rangeBetween(-49, 0)
w1 = Window.partitionBy("id", "grp").orderBy("elap").rangeBetween(-49, 0)

result = inp.withColumn(
    "count",
    F.count("lbl").over(w)
).withColumn(
    "lbl",
    F.max(F.struct("count", "lbl")).over(w1)["lbl"]
).drop("count").orderBy("id", "grp", "elap").distinct()

result.show()
#+---+---+----+---+
#| id|grp|elap|lbl|
#+---+---+----+---+
#|  1|  A|   7|  2|
#|  1|  A|  14|  3|
#|  1|  A|  35|  2|
#|  1|  A|  42|  3|
#|  1|  B|  14|  1|
#|  1|  B|  84|  2|
#|  2|  A|  14|  1|
#|  2|  A|  21|  1|
#+---+---+----+---+
Related