Correct places to cache Dataframe in Spark when unioning Dataframes

Viewed 26

I'm reading multiple files (each 200GB with couple of billions of rows) and want to process each of them, join them together, and then process again. I'm trying to figure out what are the best places to cache Dataframes. Here's my code:

def load_and_process_single_file(path):
    df = spark.read.parquet(path) # This has 1B+ rows
   
    # Some transformations -> I guess there's no need caching these?
    df = df.withColumn('a', something)
    df = df.select(...)
    df = df.where(...)

    # Calculate counts
    df = df.groupby('a').count() # This has just ~10K rows

    return df

for i in range(n):
    df_single = load_and_process_single_file(path[i])

    if(i == 0):
        df_all = df_single 
    else:
        df_all = df_all.join(df_single, df_all.a == df_single.a, 'full_outer')
        # More transformations
        df_all = df_all.withColumn(...)
        df_all = df_all.fillna(0).withColumn(...)
        df_all = df_all.select(..)

        # Does uncommenting the below line change if something should be cached?
        # df_all.show()

    # Save df_all to disk after each iteration
    df_all.write.csv(...)



Is it correct that there are no actions in load_and_process_single_file and that no caching is needed inside? Are the obvious places where .cache() should be called?

With the above code I've tried loading 7 files (n=7) and after 30 hours it still hasn't finished. Are there any performance improvements other than caching that I could use?

1 Answers

It could be worthwhile to cache df_all after a couple of iterations. This would ensure that the results are persisted in memory and disk for iterations in batches.

Related