Pyspark - filter, groupby, aggregate for different combinations of columns and functions

Viewed 223

I have a simple operation to do in Pyspark but I need to run the operation with many different parameters. It is just filter on one column, then groupby a different column, and aggregate on a third column. In Python, the function is:

def filter_gby_reduce(df, filter_col = None, filter_value = None):
  return df.filter(col(filter_col) == filter_value).groupby('ID').agg(max('Value'))

Let's say the different configurations are

func_params = spark.createDataFrame([('Day', 'Monday'), ('Month', 'January')], ['feature', 'filter_value'])

I could of course just run the functions one by one:

filter_gby_reduce(df, filter_col = 'Day', filter_value = 'Monday')
filter_gby_reduce(df, filter_col = 'Month', filter_value = 'January')

But my actual collection of parameters is much larger. Lastly, I also need to union all of the function results together into one dataframe. So is there a way in spark to write this more succinctly and in a way that will fully take advantage of parallelization?

1 Answers

One way of doing this is by generating the desired values as columns using when and max and passing these to agg. As you want the values unioned you have to unpivot the result using stack (no DataFrame API for that, so a selectExpr is used). Depending on your dataset you might get null if a filter excludes all data, these can be dropped if needed.

I'd recommend testing this vs the 'naive' approach of simply unioning a large amount of filtered dataframes.

import pyspark.sql.functions as f
func_params = [('Day', 'Monday'), ('Month', 'January')]

df = spark.createDataFrame([
    ('Monday', 'June', 1, 5), 
    ('Monday', 'January', 1, 2), 
    ('Monday', 'June', 1, 5),
    ('Monday', 'June', 2, 10)], 
    ['Day', 'Month', 'ID', 'Value'])


cols = []
for column, flt in func_params:
    name = f'{column}_{flt}'
    val = f.when(f.col(column) == flt, f.col('Value')).otherwise(None)
    cols.append(f.max(val).alias(name))

stack = f"stack({len(cols)}," + ','.join(f"'{column}_{flt}', {column}_{flt}" for column, flt in func_params) + ')'

(df
    .groupby('ID')
    .agg(*cols)
    .selectExpr('ID', stack)
    .withColumnRenamed('col0', 'param')
    .withColumnRenamed('col1', 'Value')
    .show()
)

+---+-------------+-----+                                                       
| ID|        param|Value|
+---+-------------+-----+
|  1|   Day_Monday|    5|
|  1|Month_January|    2|
|  2|   Day_Monday|   10|
|  2|Month_January| null|
+---+-------------+-----+
Related