【问题标题】:pyspark - create Top3 groups and aggregate other groups/rowspyspark - 创建 Top3 组并聚合其他组/行
【发布时间】:2017-12-30 09:18:52
【问题描述】:

我想创建一个新的数据框,其中 type 列将是基于 最高 count 的 topX。 将有一个 additional 类型 (others),它将是同一组 name 的所有 typeX 的 sum

对于 DF:

data = spark.createDataFrame([
      ("name1", "type1", 2), ("name1", "type2", 1), ("name1", "type3", 4), ("name1", "type3", 5), \
      ("name2", "type1", 6), ("name2", "type1", 7), ("name2", "type2", 8) \
    ],["name", "type", "cnt"])
    data.printSchema()

是什么:

|name  |type|cnt|
|------|-----------
|name1 |typeA|  6|
|name1 |typeX|  5|
|name1 |typeW|  3|
|name1 |typeZ|  1|
|name2 |typeA|  7|
|name2 |typeB|  2|
| .... | ... |   |  

生成的数据框(前 2 名)将是: 每个名字都有前2个值+“其他”(3组)

|name  |type|cnt|
|------|-----------
|name1 |typeA|  6|
|name1 |typeX|  5|
|name1 |other|  4|
|name2 |typeA|  7|
|name2 |typeB|  2|
|name2 |other|  0|
| .... | ... |   |  

我不确定如何跳过某个组的 X 行,然后开始聚合剩余的行。

【问题讨论】:

  • 每个名字都有重复的类型吗?您的代码似乎没有给出您正在显示的表格。

标签: python apache-spark dataframe aggregation


【解决方案1】:

我尝试使用窗口函数以及基于名称和 cnt 的行排名,然后为每个名称过滤前 2 个排名并聚合其他排名,最后合并它们。

>>> from pyspark.sql import SparkSession
>>> spark = SparkSession.builder.getOrCreate()
>>> data = spark.createDataFrame([
  ("name1", "type1", 2), ("name1", "type2", 1), ("name1", "type3", 4), ("name1", "type3", 5), \
  ("name2", "type1", 6), ("name2", "type1", 7), ("name2", "type2", 8) \
],["name", "type", "cnt"])
>>> data.show()
+-----+-----+---+
| name| type|cnt|
+-----+-----+---+
|name1|type1|  2|
|name1|type2|  1|
|name1|type3|  4|
|name1|type3|  5|
|name2|type1|  6|
|name2|type1|  7|
|name2|type2|  8|
+-----+-----+---+

>>> from pyspark.sql.window import Window
>>> from pyspark.sql.functions import rank, col,lit
>>> window = Window.partitionBy(data['name']).orderBy(data['cnt'].desc())
>>> data1 = data.select('*', rank().over(window).alias('rank'))
>>> data1.show()
+-----+-----+---+----+
| name| type|cnt|rank|
+-----+-----+---+----+
|name1|type3|  5|   1|
|name1|type3|  4|   2|
|name1|type1|  2|   3|
|name1|type2|  1|   4|
|name2|type2|  8|   1|
|name2|type1|  7|   2|
|name2|type1|  6|   3|
+-----+-----+---+----+
>>> data2 = data1.filter(data1['rank'] > 2).groupby('name').sum('cnt').select('name',lit('other').alias('type'),col('sum(cnt)').alias('cnt'))
>>> data2.show()
+-----+-----+---+
| name| type|cnt|
+-----+-----+---+
|name1|other|  3|
|name2|other|  6|
+-----+-----+---+
>>> data1.filter(data1['rank'] <=2).select('name','type','cnt').union(data2).show()
+-----+-----+---+
| name| type|cnt|
+-----+-----+---+
|name1|type3|  5|
|name1|type3|  4|
|name2|type2|  8|
|name2|type1|  7|
|name1|other|  3|
|name2|other|  6|
+-----+-----+---+

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2017-10-13
    • 1970-01-01
    • 2017-04-06
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多