How to re-assign session_id to items when we want to create another session after every null value in items?

Viewed 45

I have a pyspark dataframe-

df1 = spark.createDataFrame([
    ("s1", "i1", 0),
    ("s1", "i2", 1),
    ("s1", "i3", 2),
    ("s1", None, 3),
    ("s1", "i5", 4),

    ],
    ["session_id", "item_id", "pos"])

df1.show(truncate=False)

pos is the position or rank of the item in the session. Now I want to create new sessions without any null values in them. I want to do this by starting a new session after every null item. Basically I want to break existing sessions into multiple sessions, removing the null item_id in the process.

The expected output would like something like-

+----------+-------+---+--------------+
|session_id|item_id|pos|new_session_id|
+----------+-------+---+--------------+
|s1        |i1     |0  |          s1_0|
|s1        |i2     |1  |          s1_0|
|s1        |i3     |2  |          s1_0|
|s1        |null   |3  |          None|
|s1        |i5     |4  |          s1_4|
+----------+-------+---+--------------+

How do I achieve this?

1 Answers

Not sure about the configs of your spark job, but to prevent to use collect action to build the reference of your "new" session in Python built-in data structure, I would use built-in spark sql function to build the new session reference. Based on your example, assuming you have already sorted the data frame:

from pyspark.sql import SparkSession
from pyspark.sql import functions as func
from pyspark.sql.window import Window
from pyspark.sql.types import *

df = spark.createDataFrame(
    [("s1", "i1", 0), ("s1", "i2", 1), ("s1", "i3", 2),  ("s1", None, 3), ("s1", None, 4), ("s1", "i6", 5), ("s2", "i7", 6), ("s2", None, 7), ("s2", "i9", 8), ("s2", "i10", 9), ("s2", "i11", 10)],
    ["session_id", "item_id", "pos"]
)
df.show(20, False)

+----------+-------+---+
|session_id|item_id|pos|
+----------+-------+---+
|s1        |i1     |0  |
|s1        |i2     |1  |
|s1        |i3     |2  |
|s1        |null   |3  |
|s1        |null   |4  |
|s1        |i6     |5  |
|s2        |i7     |6  |
|s2        |null   |7  |
|s2        |i9     |8  |
|s2        |i10    |9  |
|s2        |i11    |10 |
+----------+-------+---+

Step 1: As the data is already sorted, we can use a lag function to shift the data to the next record:

df2 = df\
    .withColumn('lag_item', func.lag('item_id', 1).over(Window.partitionBy('session_id').orderBy('pos')))
df2.show(20, False)

+----------+-------+---+--------+
|session_id|item_id|pos|lag_item|
+----------+-------+---+--------+
|s1        |i1     |0  |null    |
|s1        |i2     |1  |i1      |
|s1        |i3     |2  |i2      |
|s1        |null   |3  |i3      |
|s1        |null   |4  |null    |
|s1        |i6     |5  |null    |
|s2        |i7     |6  |null    |
|s2        |null   |7  |i7      |
|s2        |i9     |8  |null    |
|s2        |i10    |9  |i9      |
|s2        |i11    |10 |i10     |
+----------+-------+---+--------+

Step 2: After using the lag function we can see if the item_id in previous record is NULL or not. Therefore , we can know the boundaries of each new session by doing the filtering and build the reference:

reference = df2\
    .filter((func.col('item_id').isNotNull())&(func.col('lag_item').isNull()))\
    .groupby('session_id')\
    .agg(func.collect_set('pos').alias('session_id_set'))
reference.show(100, False)

+----------+--------------+
|session_id|session_id_set|
+----------+--------------+
|s1        |[0, 5]        |
|s2        |[6, 8]        |
+----------+--------------+

Step 3: Join the reference back to the data and write a simple UDF to find which new session should be in:

@func.udf(returnType=IntegerType())
def udf_find_session(item_id, pos, session_id_set):
    r_val = None

    if item_id != None:
        for item in session_id_set:
            if pos >= item:
                r_val = item
            else:
                break

    return r_val

df3 = df2.select('session_id', 'item_id', 'pos')\
    .join(reference, on='session_id', how='inner')
df4 = df3.withColumn('new_session_id', udf_find_session(func.col('item_id'), func.col('pos'), func.col('session_id_set')))
df4.show(20, False)

+----------+-------+---+--------------+
|session_id|item_id|pos|new_session_id|
+----------+-------+---+--------------+
|s1        |i1     |0  |0             |
|s1        |i2     |1  |0             |
|s1        |i3     |2  |0             |
|s1        |null   |3  |null          |
|s1        |null   |4  |null          |
|s1        |i6     |5  |5             |
|s2        |i7     |6  |6             |
|s2        |null   |7  |null          |
|s2        |i9     |8  |8             |
|s2        |i10    |9  |8             |
|s2        |i11    |10 |8             |
+----------+-------+---+--------------+

The last step just concat the string you want to show in new session id.

Related