ML Engineer MasterClass (October) | 4 seats left

Spark · Filter grouped results
AmazonAmazon Analytics
Amazon · Aggregate and Analyze

Filter grouped results

Distinguish row filters from filters on aggregate values.

Step 1 of 6 · Learn

Filter a summary

After agg, filter can reference aggregate aliases. This is the DataFrame equivalent of a SQL HAVING condition. The example retains customers with total_order_value at least 150, counting all their orders.

Lesson reference: PySpark and Scala

Filter a summary

After agg, filter can reference aggregate aliases. This is the DataFrame equivalent of a SQL HAVING condition. The example retains customers with total_order_value at least 150, counting all their orders.

PySpark example

from pyspark.sql.functions import col, count, countDistinct, sum, avg, round

result = orders.groupBy("customer_id").agg(
    sum("total").alias("total_order_value")
).filter(col("total_order_value") >= 150).orderBy("customer_id")
result.show()

Scala example

import org.apache.spark.sql.functions._

val result = orders.groupBy("customer_id").agg(
    sum("total").alias("total_order_value")
).filter(col("total_order_value") >= 150).orderBy("customer_id")
result.show()

Filter before and after

A filter before groupBy removes input rows; a filter after agg removes summaries. Here only individual orders worth at least 50 contribute, then customers need an aggregate of at least 100. Customer 101 loses its 35 order before aggregation.

PySpark example

from pyspark.sql.functions import col, count, countDistinct, sum, avg, round

result = orders.filter(col("total") >= 50).groupBy("customer_id").agg(
    sum("total").alias("total_order_value")
).filter(col("total_order_value") >= 100).orderBy("customer_id")
result.show()

Scala example

import org.apache.spark.sql.functions._

val result = orders.filter(col("total") >= 50).groupBy("customer_id").agg(
    sum("total").alias("total_order_value")
).filter(col("total_order_value") >= 100).orderBy("customer_id")
result.show()

Combine summary conditions

You can filter on several aggregate aliases. The example requires at least two orders and total_order_value of at least 100. Parenthesize each comparison and combine them with & in PySpark or && in Scala.

PySpark example

from pyspark.sql.functions import col, count, countDistinct, sum, avg, round

result = orders.groupBy("customer_id").agg(
    count("*").alias("order_count"),
    sum("total").alias("total_order_value")
).filter(
    (col("order_count") >= 2) & (col("total_order_value") >= 100)
).orderBy("customer_id")
result.show()

Scala example

import org.apache.spark.sql.functions._

val result = orders.groupBy("customer_id").agg(
    count("*").alias("order_count"),
    sum("total").alias("total_order_value")
).filter(
    (col("order_count") >= 2) && (col("total_order_value") >= 100)
).orderBy("customer_id")
result.show()
example.pyPySpark
1. Read ordersOne row per order, including every status.
2. Group, then filter summariesRow filters change the inputs; summary filters change which groups remain.
3. Inspect the summaryCompare customer totals against the stated thresholds.
Source orders6 rows
order_idcustomer_idstatustotalitem_count
1001101Delivered89.52
1002102Shipped1493
1003101Cancelled351
1004103Delivered219.994
1005104Delivered49.991
1006105Shipped1202
Filter a summary
Result1 rows
customer_idtotal_order_value
103219.99
Compare the source columns with the transformed result.

The example is loaded in the editor. Run it as written, then try a small change.

Runs on the Spark backend. First startup may take a moment.

Run your code to see the result.