pyspark explode performance

Viewed 120

Background I use explode to transpose columns to rows. This works very well in general with good performance. The source dataframe (df_audit in below code) is dynamic so can contain different structure.

Problem Recently have incoming dataframe with very large number of columns (5 thousand) - the below code runs successfully but is very slow to run the line starting 'exploded'. Anyone faced similar problems? I could split up the dataframe to multiple dataframes (broken out by columns) or might there be better way? Or example code?

Example code

key_cols = ["cola", "colb", "colc"]

cols = [col for col in df_audit.columns if col not in key_cols]

exploded = explode(array([struct(lit(c).alias("key"), col(c).alias("val")) for c in cols])).alias("exploded")

df_audit =  df_audit.select(key_cols + [exploded]).select(key_cols + ["exploded.key", "exploded.val"])
3 Answers

Both lit() and col() are for some reason quite slow when used in a loop. You can try instead with arrays_zip():

exploded = explode(
    arrays_zip(split(lit(','.join(cols)), ',').alias('key'), array(cols).alias('val'))
).alias('exploded')

In my quick test on 5k columns, this runs for ~6s vs. original ~25s.

Sharing some timings for bzu's approach and OP's approach based on colaboratory notebook.

cols = ['i'+str(i) for i in range(5000)]

# OP's method
%timeit func.array(*[func.struct(func.lit(k).alias('k'), func.col(k).alias('v')) for k in cols])
# 34.7 s ± 2.84 s per loop (mean ± std. dev. of 7 runs, 1 loop each)

# bzu's method
%timeit func.arrays_zip(func.split(func.lit(','.join(cols)), ',').alias('k'), func.array(cols).alias('v'))
# 10.7 s ± 1.41 s per loop (mean ± std. dev. of 7 runs, 1 loop each)

Thank you bzu & samkart but for some reason I cannot get the new line working. I have created a simple example that doesn't work as follows if you can see something obvious I am missing.

from pyspark.sql.functions import (
    array, arrays_zip, coalesce, col, explode, lit, lower, split, struct,substring,)
from pyspark.sql.types import StringType

def process_data():
    try:
        logger.info("\ntest 1")
        df_audit = spark.createDataFrame([("1", "foo", "abc", "xyz"),("2", "bar", "def", "zab"),],["id", "label", "colx", "coly"])

        logger.info("\ntest 2")
        key_cols = ["id", "label"]
        cols = [col for col in df_audit.columns if col not in key_cols]

        logger.info("\ntest 3")
        # exploded = explode(array([struct(lit(c).alias("key"), col(c).alias("val")) for c in cols])).alias("exploded")
        exploded = explode(arrays_zip(split(lit(','.join(cols)), ',').alias('key'), array(cols).alias('val'))).alias('exploded')

        logger.info("\ntest 4")
        df_audit =  df_audit.select(key_cols + [exploded]).select(key_cols + ["exploded.key", "exploded.val"])
        df_audit.show()
    except Exception as e:
        logger.error("Error in process_audit_data: {}".format(e))
        return False
    return True

When I call process_data function I get following logged: test 1 test 2 test 3 test 4 Error in process_audit_data: No such struct field key in 0, 1. Note: it does work successfully with the commented exploded line

Many thanks

Related