Is there a way to use map_groups in agg?
# daft-dev
k
Is there a way to use map_groups in agg?
c
Currently no, because we expect that the result of agg is a single row, while map_groups is allowed to return multiple rows.
k
I see, thanks!
So if I had a group with 10 rows and I wanted to return the 5 rows with the top 5 revenues within the group for example, how would I go about returning these multiple rows?
c
hmm, couple ideas: You could use map_groups then sort the result in the udf, and use the slice method to take the top 5:
Copy code
@daft.udf(return_dtype=daft.DataType.int64())
def sort_and_limit(series: daft.Series) -> daft.Series:
    return series.sort(descending=True).slice(0, 2)


df = daft.from_pydict(
    {
        "region": ["Asia", "Asia", "Asia", "Europe", "Europe", "Europe", "America", "America", "America"],
        "store": [1, 2, 3, 4, 5, 6, 7, 8, 9],
        "revenue": [100, 200, 300, 400, 500, 600, 700, 800, 900],
    }
)

top5 = df.groupby("region").map_groups(sort_and_limit(daft.col("revenue")).alias("top_revenue")).collect()
or you could sort by your group and revenue, then manually do a limit for each group and then concat.
Copy code
df = daft.from_pydict(
    {
        "region": ["Asia", "Asia", "Asia", "Europe", "Europe", "Europe", "America", "America", "America"],
        "store": [1, 2, 3, 4, 5, 6, 7, 8, 9],
        "revenue": [100, 200, 300, 400, 500, 600, 700, 800, 900],
    }
)
# sort by region and revenue
sorted_df = df.sort(["region", "revenue"], [True, True]).collect()

all_top5_df = None
for region in sorted_df.select("region").distinct().to_pydict()["region"]:
    top5 = sorted_df.where(sorted_df["region"] == region).limit(2)
    if all_top5_df is None:
        all_top5_df = top5
    else:
        all_top5_df = all_top5_df.concat(top5)

all_top5_df.collect()
Ideally i think this should look more like a window function, e.g.
Copy code
WITH RankedRows AS (
    SELECT
        *,
        ROW_NUMBER() OVER (PARTITION BY group_column ORDER BY revenue_column DESC) AS row_num
    FROM
        your_table
)
SELECT
    *
FROM
    RankedRows
WHERE
    row_num <= 5;
but not possible in daft yet
k
Makes sense! Thanks!!