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.