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.
Join ML Engineer Interview MasterClass (October Cohort) led by FAANG Data Scientists | Just 4 seats remaining...
ML Engineer MasterClass (October) | 4 seats left
Distinguish row filters from filters on aggregate values.
Step 1 of 6 · Learn
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.
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.
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()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()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.
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()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()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.
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()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()orders6 rows| order_id | customer_id | status | total | item_count |
|---|---|---|---|---|
| 1001 | 101 | Delivered | 89.5 | 2 |
| 1002 | 102 | Shipped | 149 | 3 |
| 1003 | 101 | Cancelled | 35 | 1 |
| 1004 | 103 | Delivered | 219.99 | 4 |
| 1005 | 104 | Delivered | 49.99 | 1 |
| 1006 | 105 | Shipped | 120 | 2 |
| customer_id | total_order_value |
|---|---|
| 103 | 219.99 |
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.