使用Spark DataFrame获取分组后所有组的TopN

kwi*_*wng 9 sql scala apache-spark apache-spark-sql

我有一个Spark SQL DataFrame:

user1 item1 rating1
user1 item2 rating2
user1 item3 rating3
user2 item1 rating4
...
Run Code Online (Sandbox Code Playgroud)

如何按用户分组,然后TopN使用Scala从每个组返回项目?

使用Python的相似代码:

df.groupby("user").apply(the_func_get_TopN)
Run Code Online (Sandbox Code Playgroud)

zer*_*323 20

您可以rank按如下方式使用窗口功能

import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions.{rank, desc}

val n: Int = ???

// Window definition
val w = Window.partitionBy($"user").orderBy(desc("rating"))

// Filter
df.withColumn("rank", rank.over(w)).where($"rank" <= n)
Run Code Online (Sandbox Code Playgroud)

如果你不关心的关系,那么你可以替换rank使用row_number