Filter PySpark dataframe into a list of dataframes

Viewed 119

I have a PySpark dataframe and I want to filter based on unique values in some columns.

from pyspark.sql import SparkSession
spark_session = SparkSession.builder.enableHiveSupport().getOrCreate()

columns = ["language","users_count","apple"]
data = [("Java", 1, 0.0), ("Scala", 4, -4.0), ("Java", 1, 0.0)]

pyspark_df = spark_session.createDataFrame(data).toDF(*columns)

pandas_df = pd.DataFrame(data, columns=columns)

# Operation I want to replicate in PySpark:
column_list = ['language','users_count'] #these names and number of columns can be changed at runtime.
unique_dfs = [df for id, df in pandas_df.groupby(column_list
, as_index=False)]

Another approach that can be done is to create a column in PySpark df and put unique values (string ( language + users_count ) and later filter on those unique values to get dfs.

3 Answers

If you know exactly what data you need, you should do filter, because it is efficient in Spark.

from pyspark.sql import functions as F

df = pyspark_df.filter(
    (F.col('language') == 'Java') &
    (F.col('users_count') == 1)
)

If you REALLY need all the possible combinations of those columns as separate dataframes, you will have to run distinct (i.e. to-be-avoided shuffle) and inefficient collect

from pyspark.sql import functions as F

column_list = ['language', 'users_count']
df_dist = pyspark_df.select(column_list).distinct()
unique_dfs = []
for row in df_dist.collect():
    cond = F.lit(True)
    for c in column_list:
        cond &= (F.col(c) == row[c])
    unique_dfs.append(pyspark_df.filter(cond))

Results:

unique_dfs[0].show()
# +--------+-----------+-----+
# |language|users_count|apple|
# +--------+-----------+-----+
# |    Java|          1|  0.0|
# |    Java|          1|  0.0|
# +--------+-----------+-----+

unique_dfs[1].show()
# +--------+-----------+-----+
# |language|users_count|apple|
# +--------+-----------+-----+
# |   Scala|          4| -4.0|
# +--------+-----------+-----+

unique_dfs[0].explain()
# == Physical Plan ==
# *(1) Project [_1#158 AS language#164, _2#159L AS users_count#165L, _3#160 AS apple#166]
# +- *(1) Filter ((isnotnull(_1#158) AND isnotnull(_2#159L)) AND ((_1#158 = Java) AND (_2#159L = 1)))
#    +- *(1) Scan ExistingRDD[_1#158,_2#159L,_3#160]

Note: Here you see that Java is indexed as 0, Scala as 1, but in reality it could be opposite, you don't have determinism there, as you don't know which executor will send his data first to the driver after driver asked for data when you used collect. So, what you asked, is probably not what you truly needed.

Create rank using window function with partitioning on the columns (you want to group based on value of). Then iterate from 1 to df.count() and filter dataframe based on rank and store dataframes into list. I hope this helps!

from pyspark.sql import functions as F, Window as W

column_list = ['language', 'users_count']
unique_dfs = []
w = W.orderBy(*column_list)
df = pyspark_df.withColumn('_rank', F.dense_rank().over(w))
for i in range(1, df.agg(F.max('_rank')).head()[0] + 1):
    unique_dfs.append(df.filter(F.col('_rank') == i))

Results:

unique_dfs[0].show()
# +--------+-----------+-----+-----+
# |language|users_count|apple|_rank|
# +--------+-----------+-----+-----+
# |    Java|          1|  0.0|    1|
# |    Java|          1|  0.0|    1|
# +--------+-----------+-----+-----+

unique_dfs[1].show()
# +--------+-----------+-----+-----+
# |language|users_count|apple|_rank|
# +--------+-----------+-----+-----+
# |   Scala|          4| -4.0|    2|
# +--------+-----------+-----+-----+

unique_dfs[0].explain()
# == Physical Plan ==
# AdaptiveSparkPlan isFinalPlan=false
# +- Filter (_rank#579 = 1)
#    +- Window [dense_rank(language#571, users_count#572L) windowspecdefinition(language#571 ASC NULLS FIRST, users_count#572L ASC NULLS FIRST, specifiedwindowframe(RowFrame, unboundedpreceding$(), currentrow$())) AS _rank#579], [language#571 ASC NULLS FIRST, users_count#572L ASC NULLS FIRST]
#       +- Sort [language#571 ASC NULLS FIRST, users_count#572L ASC NULLS FIRST], false, 0
#          +- Exchange SinglePartition, ENSURE_REQUIREMENTS, [id=#975]
#             +- Project [_1#565 AS language#571, _2#566L AS users_count#572L, _3#567 AS apple#573]
#                +- Scan ExistingRDD[_1#565,_2#566L,_3#567]

I have solved this by

groups = list(pyspark_df.select(['language','users_count']).distinct().collect())

unique_campaigns_dfs = [
    pyspark_df.where((functions.col('language') == x[0]) & (functions.col('users_count') == x[1])) for x in
    groups]
Related