Better way to add column values combinations to data frame in PySpark

Viewed 247

I have a dataset that contains 3 columns, id, day, value. I need to add rows with zeros in value for all combinations of id and day.

# Simplified version of my data frame
data = [("1", "2020-04-01", 5), 
        ("2", "2020-04-01", 5), 
        ("3", "2020-04-02", 4)]
df = spark.createDataFrame(data,['id','day', 'value'])

What I have come up with is:

# Create all combinations of id and day
ids= df.select('id').distinct()
days = df.select('day').distinct()
full = ids.crossJoin(days)

# Add combinations back to df filling value with zeros
df_full = df.join(full, ['id', 'day'], 'rightouter')\
    .na.fill(value=0,subset=['value'])

Which outputs what I need:

>>> df_full.orderBy(['id','day']).show()
+---+----------+-----+
| id|       day|value|
+---+----------+-----+
|  1|2020-04-01|    5|
|  1|2020-04-02|    0|
|  2|2020-04-01|    5|
|  2|2020-04-02|    0|
|  3|2020-04-01|    0|
|  3|2020-04-02|    4|
+---+----------+-----+

The problem is that both of these operations a very computationally expensive. When I'm running it with my full data, it gives me a job that an order of magnitude larger than something that usually takes a couple of hours to run.

Is there a more efficient way of doing this? Or is there something I'm missing?

2 Answers

That's the way I would implement. Just a point, both dataframes must have the same schema, otherwise stack function will raise an error

import pyspark.sql.functions as f


# Simplified version of my data frame
data = [("1", "2020-04-01", 5), 
        ("2", "2020-04-01", 5), 
        ("3", "2020-04-02", 4)]
df = spark.createDataFrame(data, ['id', 'day', 'value'])

# Creating a dataframe with all distinct days
df_days = df.select(f.col('day').alias('r_day')).distinct()

# Self Join to find all combinations
df_final = df.join(df_days, on=df['day'] != df_days['r_day'])
# +---+----------+-----+----------+
# | id|       day|value|     r_day|
# +---+----------+-----+----------+
# |  1|2020-04-01|    5|2020-04-02|
# |  2|2020-04-01|    5|2020-04-02|
# |  3|2020-04-02|    4|2020-04-01|
# +---+----------+-----+----------+

# Unpivot dataframe
df_final = df_final.select('id', f.expr('stack(2, day, value, r_day, cast(0 as bigint)) as (day, value)'))
df_final.orderBy('id', 'day').show()

Output:

+---+----------+-----+
| id|       day|value|
+---+----------+-----+
|  1|2020-04-01|    5|
|  1|2020-04-02|    0|
|  2|2020-04-01|    5|
|  2|2020-04-02|    0|
|  3|2020-04-01|    0|
|  3|2020-04-02|    4|
+---+----------+-----+

Something like this. You could, I keep the first row separate since it's more clear what happens. You could add it to the "main loop" though.

data = [
    ("1", date(2020, 4, 1), 5),
    ("2", date(2020, 4, 2), 5),
    ("3", date(2020, 4, 3), 5),
    ("1", date(2020, 4, 3), 5),
]


df = spark.createDataFrame(data, ["id", "date", "value"])

row_dates = df.select("date").distinct().collect()

dates = [item.asDict()["date"] for item in row_dates]


def map_row(dates: List[date]) -> Callable[[Iterator[Row]], Iterator[Row]]:
    dates.sort()

    def inner(partition):
        last_row = None

        for row in partition:
            # fill in missing dates for first row in partition
            if last_row is None:
                for day in dates:
                    if day < row.date:
                        yield Row(row.id, day, 0)
                    else:
                        # set current row as last row, yield current row and break out of the loop
                        last_row = row
                        yield row
                        break
            else:
                # if current row has same id as last row
                if last_row.id == row.id:
                    # yield dates between last and current
                    for day in dates:
                        if day > last_row.date and day < row.date:
                            yield Row(row.id, day, 0)
                    
                    # set current as last and yield current
                    last_row = row
                    yield row

                else:
                    # if current row is new id
                    for day in dates:
                        # run potential remaining dates for last_row.id
                        if day > last_row.date:
                            yield Row(last_row.id, day, 0)

                    for day in dates:
                        # fill in missing dates before row.date
                        if day < row.date:
                            yield Row(row.id, day, 0)                    
                        else:
                            # unt so weiter
                            last_row = row
                            yield row
                            break

    return inner


rdd = (
    df.repartition(1, "id")
    .sortWithinPartitions("id", "date")
    .rdd.mapPartitions(map_row(dates))
)
new_df = spark.createDataFrame(rdd)
new_df.show(10, False)
Related