Iterating through tables with map function and function that queries other dataframes

Viewed 35

I have two tables: Table A

|Group ID  | User ids in group|
| -------- | -------------- |
| 11       | [45,46,47,48]  |
| 20       | [49,10,11,12]  |
| 31       | [55,7,48,43]   |

and Table B:

| User ids| Related Id     |
| ------- | -------------- |
| 1       | [5,6,7,8]      |
| 2       | [6, 9, 10,11]  |
| 3       | [1, 2, 5, 7]   |

And I have a reference table that has the info: Reference table:

| User ids | Group ID |
| -------- | -------------- |
| 1        | 11             |
| 2        | 20             |
| 3        | 31             |

This is just a minimal sample, I have this situation with millions of rows on each table. I am trying to use pyspark (or sql but I haven't figured out a way to do it here) to iterate through User ids column in the reference table and get the intersection between the lists of User ids in group from Table A and Related Id from Table B.

So in the end, I would like to have a table of the form:

| User ids | Intersection   |
| -------- | -------------- |
| 2        | [10, 11]       |
| 3        | [7]            |

In Pyspark Id have a function of the form:

def test_function(user_id, ref_df, tableB_df, tableA_df):
    group_id = int(ref_df.filter(ref_df.userID == user_id).collect()[0][1])
    group_list = tableA_df.filter(tableA_df.groupID == group_id)
    related_id_list = tableB_df.filter(tableB_df.userID == user_id)
    
    return group_list.intsersection(related_id_list)


abc = ref_df.rdd.map(lambda x: test_function(x, ref_df, tableB_df, tableA_df))

However, when I run this function I am getting the following error:

An error was encountered: Could not serialize object: TypeError: can't pickle _thread.RLock objects

Can anyone give any suggestion on how to solve this problem? Or how to modify my approach to solve this problem? Since my table has millions of rows, I want to use pyspark as best as possible to make use of the parallelization abilities as much as possible. Thanks for all your help.

1 Answers

You first join the reference table with table a on Group ID, and join the resulting table with table b on User ids. This will give you a dataframe that looks like this:

+--------+--------+-----------------+--------------+
|User ids|Group ID|User ids in group|    Related Id|
+--------+--------+-----------------+--------------+
|       1|      11| [45, 46, 47, 48]|  [5, 6, 7, 8]|
|       2|      20| [49, 10, 11, 12]|[6, 9, 10, 11]|
|       3|      31|  [55, 7, 48, 43]|  [1, 2, 5, 7]|
+--------+--------+-----------------+--------------+

Then you perform an intersection on column User ids in group and Related Id. This gives you the columns you want, but you need to filter rows where the intersection is empty.

The code snippet below does all of that in pyspark:

import pyspark.sql.functions as F
# Init example tables
table_a = spark.createDataFrame(
    [(11, [45, 46, 47, 48]), (20, [49, 10, 11, 12]), (31, [55, 7, 48, 43])],
    ["Group ID", "User ids in group"],
)
table_b = spark.createDataFrame(
    [(1, [5, 6, 7, 8]), (2, [6, 9, 10, 11]), (3, [1, 2, 5, 7])],
    ["User ids", "Related Id"],
)
reference_table = spark.createDataFrame(
    [(1, 11), (2, 20), (3, 31)], ["User ids", "Group ID"]
)
# Relevant code
joined_df = reference_table.join(table_a, on="Group ID").join(table_b, on="User ids")
intersected_df = joined_df.withColumn("Intersection", F.array_intersect("User ids in group", "Related Id"))
intersected_df.select("User ids", "Intersection").filter(F.size("Intersection") > 0).show()

output:

+--------+------------+
|User ids|Intersection|
+--------+------------+
|       2|    [10, 11]|
|       3|         [7]|
+--------+------------+
Related