Pyspark generating a segment array using values and thresholds in 2 dataframes

Viewed 333

I need to generate a segment array using the segment values and their thresholds in 2 different datasets. Is there a simple way to do this in pyspark or hive sql?

Segment values dataset:

--------------------------------------------------
| user_id   | seg1  | seg2  | seg3 | seg4 | seg5 |
------------------------------------------------
| 100       |   90  |  20   |   76 |  100 |  30  |
| 200       |   56  |  15   |   67 |  99  |  25  |
| 300       |   87  |  38   |   45 |  97  |  40  |
--------------------------------------------------

segment threshold dataset:

---------------------------
|seg_name | seg_threshold |
---------------------------
|  seg1   |  83           |
|  seg2   |  25           |
|  seg3   |  60           |
|  seg4   |  98           |
|  seg5   |  35           |
---------------------------

If the value for a segment is higher than the threshold, the user should be considered to be part of the segment. Segment array for that user should include the segment names(column headers).

Expected output:

-------------------------------------
| user_id| segment_array            |
-------------------------------------
| 100    | [seg1, seg3, seg4] |
| 200    | [seg3, seg4]             |
| 300    | [seg1, seg2, seg5]       |
-------------------------------------

Please note that this is just an indicative dataset. I have several hundreds of segments like these.

Thank you for your help!

2 Answers

A few hundered threshold entries could be broadcasted. The check if a value is above or below the threshold can then be done in an UDF:

#broadcast the threshold data
thresholdDf = ...
thresholdMap = thresholdDf.rdd.collectAsMap()
thresholds = spark.sparkContext.broadcast(thresholdMap)

userDf = ...

#add a new column to the user dataframe that contains a struct with the column 
#names and their respective values. This column will be used to call the udf
user2Df = userDf.withColumn("all_cols", F.struct([F.struct(F.lit(x),userDf[x]) \
    for x in userDf.columns]))

#create the udf
def calc_segments(row):
    return [col.col1 for col in row \
        if thresholds.value.get(col.col1) != None \
        if int(thresholds.value[col.col1]) < int(col[col.col1])]
segment_udf = F.udf(calc_segments, T.ArrayType(T.StringType()))

#call the udf and drop the intermediate column
user2Df.withColumn("segment_array", segment_udf(user2Df.all_cols)) \
    .drop("all_cols").show(truncate=False)

My result is

+-------+----+----+----+----+----+------------------+
|user_id|seg1|seg2|seg3|seg4|seg5|segment_array     |
+-------+----+----+----+----+----+------------------+
|100    |90  |20  |76  |100 |30  |[seg1, seg3, seg4]|
|200    |56  |15  |67  |99  |25  |[seg3, seg4]      |
|300    |87  |38  |45  |97  |40  |[seg1, seg2, seg5]|
+-------+----+----+----+----+----+------------------+

This result is slightly different from the expected result. Maybe there is an issue with the test data.

@werner's solution is completely valid.

There is a way to do this without a udf, in pure spark-sql.

Prepare the data frames:

from pyspark.sql import Row

spark.createDataFrame([
  Row(user_id=100, seg1=90, seg2=20, seg3=76, seg4=100, seg5=30), 
  Row(user_id=200, seg1=56, seg2=15, seg3=67, seg4=99, seg5=25), 
  Row(user_id=300, seg1=87, seg2=38, seg3=45, seg4=97, seg5=40)]).createOrReplaceTempView("data")

spark.createDataFrame([
  Row(seg_name = 'seg1', seg_threshold = 83),
  Row(seg_name = 'seg2', seg_threshold = 25),
  Row(seg_name = 'seg3', seg_threshold = 60),
  Row(seg_name = 'seg4', seg_threshold = 98),
  Row(seg_name = 'seg5', seg_threshold = 35)
]).createOrReplaceTempView("thr")

Now, you can perform an 'unpivot' operation using a marginal but very useful function called stack:

spark.sql("""
WITH data_eva 
     AS (SELECT user_id, 
                Stack(5, 'seg1', seg1, 'seg2', seg2, 'seg3', seg3, 'seg4', seg4, 'seg5', seg5) 
         FROM   data) 
SELECT user_id, 
       Collect_list(col0) 
FROM   data_eva 
       JOIN thr 
         ON data_eva.col0 = thr.seg_name 
WHERE  col1 > seg_threshold 
GROUP  BY user_id 
 """).show()

And this is the output:

+-------+------------------+
|user_id|collect_list(col0)|
+-------+------------------+
|    100|[seg4, seg1, seg3]|
|    200|      [seg4, seg3]|
|    300|[seg2, seg1, seg5]|
+-------+------------------+

You mentioned you have hundreds of segments. You can easily generate the expression inside the stack function with a loop.

This technique is very useful to have in your spark toolbox.

Related